From 100e5a8239396db0153ae4ec89da61485c2a6769 Mon Sep 17 00:00:00 2001
From: sneeker
Date: Fri, 18 Sep 2026 15:14:34 +0000
Subject: [PATCH] initial
---
.gitignore | 10 +
README.md | 14 +
bin/dune | 4 +
bin/main.ml | 289 +
dune-project | 2 +
examples/anbn-paper.corpus | 7 +
examples/palindrome-paper.grammar | 4 +
experiments/a2n-final.svg | 106 +
experiments/a2n-final.txt | 4 +
experiments/abn-final.svg | 137 +
experiments/abn-final.txt | 5 +
experiments/addition-final.svg | 210 +
experiments/addition-final.txt | 8 +
experiments/anbn-final.svg | 144 +
experiments/anbn-final.txt | 5 +
experiments/basic_english-final.svg | 690 ++
experiments/basic_english-final.txt | 23 +
experiments/expr-final.svg | 230 +
experiments/expr-final.txt | 9 +
experiments/nested_anbn-final.svg | 142 +
experiments/nested_anbn-final.txt | 5 +
experiments/nested_parens-final.svg | 158 +
experiments/nested_parens-final.txt | 6 +
experiments/palindrome-final.svg | 176 +
experiments/palindrome-final.txt | 6 +
experiments/parens-final.svg | 160 +
experiments/parens-final.txt | 6 +
experiments/results.tsv | 12 +
experiments/shape-final.svg | 233 +
experiments/shape-final.txt | 8 +
experiments/steps/a2n-step001-after.svg | 166 +
experiments/steps/a2n-step001-before.svg | 175 +
experiments/steps/a2n-step002-after.svg | 106 +
experiments/steps/a2n-step002-before.svg | 166 +
experiments/steps/abn-step001-after.svg | 197 +
experiments/steps/abn-step001-before.svg | 206 +
experiments/steps/abn-step002-after.svg | 137 +
experiments/steps/abn-step002-before.svg | 197 +
experiments/steps/addition-step001-after.svg | 360 +
experiments/steps/addition-step001-before.svg | 592 ++
experiments/steps/addition-step002-after.svg | 210 +
experiments/steps/addition-step002-before.svg | 360 +
experiments/steps/anbn-step001-after.svg | 218 +
experiments/steps/anbn-step001-before.svg | 206 +
experiments/steps/anbn-step002-after.svg | 144 +
experiments/steps/anbn-step002-before.svg | 218 +
.../steps/basic_english-step001-after.svg | 1026 ++
.../steps/basic_english-step001-before.svg | 1311 +++
.../steps/basic_english-step002-after.svg | 938 ++
.../steps/basic_english-step002-before.svg | 1026 ++
.../steps/basic_english-step003-after.svg | 850 ++
.../steps/basic_english-step003-before.svg | 938 ++
.../steps/basic_english-step004-after.svg | 810 ++
.../steps/basic_english-step004-before.svg | 850 ++
.../steps/basic_english-step005-after.svg | 770 ++
.../steps/basic_english-step005-before.svg | 810 ++
.../steps/basic_english-step006-after.svg | 730 ++
.../steps/basic_english-step006-before.svg | 770 ++
.../steps/basic_english-step007-after.svg | 690 ++
.../steps/basic_english-step007-before.svg | 730 ++
experiments/steps/expr-step001-after.svg | 290 +
experiments/steps/expr-step001-before.svg | 1214 +++
experiments/steps/expr-step002-after.svg | 256 +
experiments/steps/expr-step002-before.svg | 290 +
experiments/steps/expr-step003-after.svg | 230 +
experiments/steps/expr-step003-before.svg | 256 +
.../steps/nested_anbn-step001-after.svg | 142 +
.../steps/nested_anbn-step001-before.svg | 278 +
.../steps/nested_parens-step001-after.svg | 158 +
.../steps/nested_parens-step001-before.svg | 829 ++
.../steps/palindrome-step001-after.svg | 176 +
.../steps/palindrome-step001-before.svg | 580 +
experiments/steps/parens-step001-after.svg | 316 +
experiments/steps/parens-step001-before.svg | 346 +
experiments/steps/parens-step002-after.svg | 160 +
experiments/steps/parens-step002-before.svg | 316 +
experiments/steps/shape-step001-after.svg | 698 ++
experiments/steps/shape-step001-before.svg | 742 ++
experiments/steps/shape-step002-after.svg | 295 +
experiments/steps/shape-step002-before.svg | 698 ++
experiments/steps/shape-step003-after.svg | 307 +
experiments/steps/shape-step003-before.svg | 295 +
experiments/steps/shape-step004-after.svg | 233 +
experiments/steps/shape-step004-before.svg | 307 +
lib/corpus.ml | 67 +
lib/dot.ml | 143 +
lib/dune | 4 +
lib/experiment.ml | 282 +
lib/grammar.ml | 522 +
lib/parse.ml | 355 +
lib/scoring.ml | 278 +
lib/search.ml | 250 +
lib/transform.ml | 193 +
meta/paper-and-induced.svg | 9380 +++++++++++++++++
test/dune | 4 +
test/test_scfg.ml | 317 +
96 files changed, 39927 insertions(+)
create mode 100644 .gitignore
create mode 100644 README.md
create mode 100644 bin/dune
create mode 100644 bin/main.ml
create mode 100644 dune-project
create mode 100644 examples/anbn-paper.corpus
create mode 100644 examples/palindrome-paper.grammar
create mode 100644 experiments/a2n-final.svg
create mode 100644 experiments/a2n-final.txt
create mode 100644 experiments/abn-final.svg
create mode 100644 experiments/abn-final.txt
create mode 100644 experiments/addition-final.svg
create mode 100644 experiments/addition-final.txt
create mode 100644 experiments/anbn-final.svg
create mode 100644 experiments/anbn-final.txt
create mode 100644 experiments/basic_english-final.svg
create mode 100644 experiments/basic_english-final.txt
create mode 100644 experiments/expr-final.svg
create mode 100644 experiments/expr-final.txt
create mode 100644 experiments/nested_anbn-final.svg
create mode 100644 experiments/nested_anbn-final.txt
create mode 100644 experiments/nested_parens-final.svg
create mode 100644 experiments/nested_parens-final.txt
create mode 100644 experiments/palindrome-final.svg
create mode 100644 experiments/palindrome-final.txt
create mode 100644 experiments/parens-final.svg
create mode 100644 experiments/parens-final.txt
create mode 100644 experiments/results.tsv
create mode 100644 experiments/shape-final.svg
create mode 100644 experiments/shape-final.txt
create mode 100644 experiments/steps/a2n-step001-after.svg
create mode 100644 experiments/steps/a2n-step001-before.svg
create mode 100644 experiments/steps/a2n-step002-after.svg
create mode 100644 experiments/steps/a2n-step002-before.svg
create mode 100644 experiments/steps/abn-step001-after.svg
create mode 100644 experiments/steps/abn-step001-before.svg
create mode 100644 experiments/steps/abn-step002-after.svg
create mode 100644 experiments/steps/abn-step002-before.svg
create mode 100644 experiments/steps/addition-step001-after.svg
create mode 100644 experiments/steps/addition-step001-before.svg
create mode 100644 experiments/steps/addition-step002-after.svg
create mode 100644 experiments/steps/addition-step002-before.svg
create mode 100644 experiments/steps/anbn-step001-after.svg
create mode 100644 experiments/steps/anbn-step001-before.svg
create mode 100644 experiments/steps/anbn-step002-after.svg
create mode 100644 experiments/steps/anbn-step002-before.svg
create mode 100644 experiments/steps/basic_english-step001-after.svg
create mode 100644 experiments/steps/basic_english-step001-before.svg
create mode 100644 experiments/steps/basic_english-step002-after.svg
create mode 100644 experiments/steps/basic_english-step002-before.svg
create mode 100644 experiments/steps/basic_english-step003-after.svg
create mode 100644 experiments/steps/basic_english-step003-before.svg
create mode 100644 experiments/steps/basic_english-step004-after.svg
create mode 100644 experiments/steps/basic_english-step004-before.svg
create mode 100644 experiments/steps/basic_english-step005-after.svg
create mode 100644 experiments/steps/basic_english-step005-before.svg
create mode 100644 experiments/steps/basic_english-step006-after.svg
create mode 100644 experiments/steps/basic_english-step006-before.svg
create mode 100644 experiments/steps/basic_english-step007-after.svg
create mode 100644 experiments/steps/basic_english-step007-before.svg
create mode 100644 experiments/steps/expr-step001-after.svg
create mode 100644 experiments/steps/expr-step001-before.svg
create mode 100644 experiments/steps/expr-step002-after.svg
create mode 100644 experiments/steps/expr-step002-before.svg
create mode 100644 experiments/steps/expr-step003-after.svg
create mode 100644 experiments/steps/expr-step003-before.svg
create mode 100644 experiments/steps/nested_anbn-step001-after.svg
create mode 100644 experiments/steps/nested_anbn-step001-before.svg
create mode 100644 experiments/steps/nested_parens-step001-after.svg
create mode 100644 experiments/steps/nested_parens-step001-before.svg
create mode 100644 experiments/steps/palindrome-step001-after.svg
create mode 100644 experiments/steps/palindrome-step001-before.svg
create mode 100644 experiments/steps/parens-step001-after.svg
create mode 100644 experiments/steps/parens-step001-before.svg
create mode 100644 experiments/steps/parens-step002-after.svg
create mode 100644 experiments/steps/parens-step002-before.svg
create mode 100644 experiments/steps/shape-step001-after.svg
create mode 100644 experiments/steps/shape-step001-before.svg
create mode 100644 experiments/steps/shape-step002-after.svg
create mode 100644 experiments/steps/shape-step002-before.svg
create mode 100644 experiments/steps/shape-step003-after.svg
create mode 100644 experiments/steps/shape-step003-before.svg
create mode 100644 experiments/steps/shape-step004-after.svg
create mode 100644 experiments/steps/shape-step004-before.svg
create mode 100644 lib/corpus.ml
create mode 100644 lib/dot.ml
create mode 100644 lib/dune
create mode 100644 lib/experiment.ml
create mode 100644 lib/grammar.ml
create mode 100644 lib/parse.ml
create mode 100644 lib/scoring.ml
create mode 100644 lib/search.ml
create mode 100644 lib/transform.ml
create mode 100644 meta/paper-and-induced.svg
create mode 100644 test/dune
create mode 100644 test/test_scfg.ml
diff --git a/.gitignore b/.gitignore
new file mode 100644
index 0000000..64c6308
--- /dev/null
+++ b/.gitignore
@@ -0,0 +1,10 @@
+_build/
+_opam/
+*.install
+.merlin
+out/
+book/
+*.cmi
+*.cmo
+*.cmx
+*.o
diff --git a/README.md b/README.md
new file mode 100644
index 0000000..4af68f2
--- /dev/null
+++ b/README.md
@@ -0,0 +1,14 @@
+An implementation of stochastic context free grammar induction, following Stolcke and Omohundro's
+[Inducing Probabilistic Grammars by Bayesian Model Merging](https://arxiv.org/abs/cmp-lg/9409010) (ICGI 1994).
+
+
+
+
+
+We start from the most specific grammar the data permits: every sample contributes its own production, and every terminal that occurs gets a corresponding nonterminal. At this stage there's effectively no sharing between samples, so the grammar is just memorising the corpus rather than generalising beyond it.
+
+From there, we generalise with two operators, merging and chunking. Merging takes a pair of nonterminals and folds them into a single nonterminal containing the union of their productions; chunking replaces a contiguous sequence of symbols with a fresh nonterminal. Chunking doesn't itself change the language the grammar generates, but it changes the internal structure in a way that can expose useful merges which weren't previously available.
+
+We rank candidate grammars by the posterior `P(M | X) ∝ P(M) P(X | M)`, where the prior is a description length and the likelihood integrates over the production probabilities under symmetric Dirichlet priors. The scoring therefore accounts for uncertainty in the production probabilities, rather than relying on a single fitted parameterisation.
+
+We explore the resulting grammar space using either beam or best first search, then fit the parameters by expectation maximisation once the grammar structure is fixed. Parsing is a generalised CYK or inside computation over spans, which lets the grammar be scored directly without an intermediate conversion into Chomsky normal form.
diff --git a/bin/dune b/bin/dune
new file mode 100644
index 0000000..388bec0
--- /dev/null
+++ b/bin/dune
@@ -0,0 +1,4 @@
+(executable
+ (name main)
+ (flags (:standard -warn-error -a))
+ (libraries scfg))
diff --git a/bin/main.ml b/bin/main.ml
new file mode 100644
index 0000000..7e952e5
--- /dev/null
+++ b/bin/main.ml
@@ -0,0 +1,289 @@
+open Scfg
+
+let read_file path =
+ let ic = open_in path in
+ let n = in_channel_length ic in
+ let s = really_input_string ic n in
+ close_in ic;
+ s
+
+let mkdir_p dir = ignore (Sys.command ("mkdir -p " ^ Filename.quote dir))
+
+let parse_args rest =
+ let rec go opts = function
+ | [] -> opts
+ | x :: v :: tl
+ when String.length x > 2
+ && String.sub x 0 2 = "--"
+ && String.length v > 0
+ && v.[0] <> '-' ->
+ go ((String.sub x 2 (String.length x - 2), v) :: opts) tl
+ | x :: tl when String.length x > 2 && String.sub x 0 2 = "--" ->
+ go ((String.sub x 2 (String.length x - 2), "true") :: opts) tl
+ | _ :: tl -> go opts tl
+ in
+ go [] rest
+
+let opt opts k d = match List.assoc_opt k opts with Some v -> v | None -> d
+
+let float_opt opts k d =
+ match List.assoc_opt k opts with
+ | Some v -> (match float_of_string_opt v with Some f -> f | None -> d)
+ | None -> d
+
+let int_opt opts k d =
+ match List.assoc_opt k opts with
+ | Some v -> (match int_of_string_opt v with Some i -> i | None -> d)
+ | None -> d
+
+let flag opts k = List.mem_assoc k opts
+
+let config_of_opts opts =
+ {
+ Scoring.default_config with
+ alpha = float_opt opts "alpha" Scoring.default_config.alpha;
+ prior_weight = float_opt opts "prior" Scoring.default_config.prior_weight;
+ em_iters = int_opt opts "em-iters" Scoring.default_config.em_iters;
+ max_chunk = int_opt opts "max-chunk" Scoring.default_config.max_chunk;
+ mode =
+ (match opt opts "mode" "marginal" with
+ | "ml" -> Scoring.Maximum_likelihood
+ | "variational" -> Scoring.Variational
+ | _ -> Scoring.Marginal);
+ chunk_occurrence =
+ (if opt opts "chunk-occurrence" "all" = "first" then
+ Scoring.First_occurrence
+ else Scoring.All_occurrences);
+ }
+
+let usage () =
+ print_endline
+ "SCFG induction by Bayesian model merging (Stolcke & Omohundro 1994)\n\n\
+ Usage: scfg-induction [options]\n\n\
+ Commands:\n\
+ \ induce --corpus FILE [--out DIR] [--beam N] [--max-steps N]\n\
+ \ [--patience N] [--alpha A] [--prior W] [--max-chunk N]\n\
+ \ [--em-iters N] [--mode marginal|ml|variational]\n\
+ \ [--chunk-occurrence all|first] [--render]\n\
+ \ [--search beam|bestfirst] [--frontier N] [--max-expansions N]\n\
+ \ score --grammar FILE --corpus FILE [--alpha A] [--prior W]\n\
+ \ parse --grammar FILE --input \"a b c\" [--dot FILE]\n\
+ \ dot --grammar FILE --out FILE [--svg FILE] [--title T]\n\
+ \ experiments [--out DIR] [--seed N] [--beam N] [--max-steps N]\n\
+ \ [--patience N] [--max-len N] [--only NAME] [--quick]\n\n\
+ Corpus format: one sequence per line, whitespace-separated tokens.\n\
+ An optional \"count:\" prefix sets the sample count, e.g. \"3: a b c\".\n\
+ Grammar format: \"S -> A B [0.5]\"; uppercase-initial symbols are\n\
+ nonterminals, other symbols are terminals."
+
+let cmd_induce opts =
+ let corpus_path = opt opts "corpus" "" in
+ if corpus_path = "" then (prerr_endline "induce: --corpus is required"; exit 2);
+ let out_dir = opt opts "out" "out" in
+ mkdir_p out_dir;
+ let steps_dir = Filename.concat out_dir "steps" in
+ mkdir_p steps_dir;
+ let render = flag opts "render" in
+ let config = config_of_opts opts in
+ let beam = int_opt opts "beam" 4 in
+ let max_steps = int_opt opts "max-steps" 25 in
+ let patience = int_opt opts "patience" 5 in
+ let corpus = Corpus.of_file corpus_path in
+ let initial =
+ Grammar.initial_grammar
+ (List.map (fun s -> (s.Corpus.tokens, s.Corpus.count)) corpus)
+ in
+ Printf.printf "corpus: %d distinct samples, total count %g\n"
+ (List.length corpus) (Corpus.total_count corpus);
+ Printf.printf "initial grammar: %d nonterminals, %d productions\n"
+ (List.length (Grammar.nonterminals initial))
+ (List.length initial.Grammar.productions);
+ let result =
+ if opt opts "search" "beam" = "bestfirst" then
+ Search.run_best_first ~config
+ ~frontier_cap:(int_opt opts "frontier" 200)
+ ~max_expansions:(int_opt opts "max-expansions" 300)
+ ~patience initial corpus
+ else Search.run ~config ~beam_width:beam ~max_steps ~patience initial corpus
+ in
+ List.iteri
+ (fun i (s : Search.step) ->
+ let ch = Transform.change_of s.Search.parent s.Search.transform s.Search.result in
+ let base =
+ Filename.concat steps_dir (Printf.sprintf "step%03d" (i + 1))
+ in
+ let dot_before =
+ Dot.dot_grammar ~highlight_nts:ch.Transform.before_nts
+ ~highlight_prods:ch.Transform.before_prods
+ ~title:
+ (Printf.sprintf "step %d before: %s" (i + 1)
+ (Transform.describe s.Search.transform))
+ s.Search.parent
+ in
+ let dot_after =
+ Dot.dot_grammar ~highlight_nts:ch.Transform.after_nts
+ ~highlight_prods:ch.Transform.after_prods
+ ~title:
+ (Printf.sprintf "step %d after: %s" (i + 1)
+ (Transform.describe s.Search.transform))
+ s.Search.result
+ in
+ let svg_before = if render then Some (base ^ "-before.svg") else None in
+ let svg_after = if render then Some (base ^ "-after.svg") else None in
+ ignore
+ (Dot.render ?out_svg:svg_before ~out_dot:(base ^ "-before.dot") dot_before);
+ ignore
+ (Dot.render ?out_svg:svg_after ~out_dot:(base ^ "-after.dot") dot_after);
+ Printf.printf " accepted step %d: %s score %.4f -> %.4f\n" (i + 1)
+ (Transform.describe s.Search.transform)
+ s.Search.parent_score s.Search.score)
+ result.Search.steps;
+
+ let details = Scoring.details ~config result.Search.best corpus in
+ Dot.write_file (Filename.concat out_dir "grammar-final.dot")
+ (Dot.dot_grammar ~title:"induced grammar" result.Search.best);
+ if render then
+ ignore
+ (Dot.render ~out_svg:(Filename.concat out_dir "grammar-final.svg")
+ ~out_dot:(Filename.concat out_dir "grammar-final.dot")
+ (Dot.dot_grammar ~title:"induced grammar" result.Search.best));
+ Dot.write_file (Filename.concat out_dir "grammar-initial.dot")
+ (Dot.dot_grammar ~title:"initial grammar" initial);
+ Dot.write_file (Filename.concat out_dir "grammar-final.txt")
+ (Grammar.to_string result.Search.best);
+ Dot.write_file (Filename.concat out_dir "search-steps.tsv")
+ (Search.steps_to_tsv result);
+ Dot.write_file (Filename.concat out_dir "search-proposals.tsv")
+ (Search.proposals_to_tsv result);
+ Printf.printf "\nsearch: %d iterations, %d accepted steps\n"
+ result.Search.iterations (List.length result.Search.steps);
+ Printf.printf "posterior: %.4f -> %.4f\n" result.Search.initial_score
+ result.Search.best_score;
+ Printf.printf "structure prior %.4f, marginal loglik %.4f, dl %.1f bits\n"
+ details.Scoring.prior details.Scoring.marginal details.Scoring.dl_bits;
+ Printf.printf "final grammar: %d nonterminals, %d productions\n"
+ details.Scoring.num_nonterminals details.Scoring.num_productions;
+ print_endline "\nfinal grammar:";
+ print_endline (Grammar.to_string result.Search.best);
+ Printf.printf "\nwrote %s\n" out_dir
+
+let cmd_score opts =
+ let grammar_path = opt opts "grammar" "" in
+ let corpus_path = opt opts "corpus" "" in
+ if grammar_path = "" || corpus_path = "" then
+ (prerr_endline "score: --grammar and --corpus are required"; exit 2);
+ let config = config_of_opts opts in
+ let g = Grammar.of_string (read_file grammar_path) in
+ let corpus = Corpus.of_file corpus_path in
+ let d = Scoring.details ~config g corpus in
+ Printf.printf "posterior %.6f\n" d.Scoring.posterior;
+ Printf.printf "structural prior %.6f\n" d.Scoring.prior;
+ Printf.printf "marginal loglik %.6f\n" d.Scoring.marginal;
+ Printf.printf "ML corpus loglik %.6f\n" d.Scoring.ml_loglik;
+ Printf.printf "description length %.2f bits\n" d.Scoring.dl_bits;
+ Printf.printf "nonterminals %d\n" d.Scoring.num_nonterminals;
+ Printf.printf "productions %d\n" d.Scoring.num_productions
+
+let cmd_parse opts =
+ let grammar_path = opt opts "grammar" "" in
+ let input = opt opts "input" "" in
+ if grammar_path = "" || input = "" then
+ (prerr_endline "parse: --grammar and --input are required"; exit 2);
+ let g = Grammar.normalize (Grammar.of_string (read_file grammar_path)) in
+ let tokens =
+ String.split_on_char ' ' (String.trim input)
+ |> List.filter (fun s -> String.length s > 0)
+ in
+ let prep = Parse.prepare g in
+ let p = Parse.inside prep tokens in
+ let vp, tree = Parse.viterbi prep tokens in
+ Printf.printf "tokens: %s\n" (String.concat " " tokens);
+ Printf.printf "P(string): %.10g\n" p;
+ Printf.printf "best parse: %.10g\n" vp;
+ (match tree with
+ | Some t -> Printf.printf "%s\n" (Parse.tree_to_string t)
+ | None -> print_endline "(no parse)");
+ (match List.assoc_opt "dot" opts with
+ | Some path -> (
+ match tree with
+ | Some t ->
+ Dot.write_file path (Dot.dot_tree ~title:"Viterbi parse" t);
+ Printf.printf "wrote %s\n" path
+ | None -> ())
+ | None -> ())
+
+let cmd_dot opts =
+ let grammar_path = opt opts "grammar" "" in
+ let out = opt opts "out" "" in
+ if grammar_path = "" || out = "" then
+ (prerr_endline "dot: --grammar and --out are required"; exit 2);
+ let g = Grammar.normalize (Grammar.of_string (read_file grammar_path)) in
+ let title = opt opts "title" "grammar" in
+ let svg = List.assoc_opt "svg" opts in
+ let res =
+ Dot.render ?out_svg:svg ~out_dot:out (Dot.dot_grammar ~title g)
+ in
+ Printf.printf "wrote %s%s\n" out
+ (match res with Some s -> " and " ^ s | None -> "")
+
+let cmd_experiments opts =
+ let out_dir = opt opts "out" "experiments" in
+ mkdir_p out_dir;
+ mkdir_p (Filename.concat out_dir "steps");
+ let render = flag opts "render" in
+ let quick = flag opts "quick" in
+ let config = config_of_opts opts in
+ let seed = int_opt opts "seed" 1 in
+ let beam = int_opt opts "beam" 3 in
+ let max_steps = int_opt opts "max-steps" (if quick then 12 else 25) in
+ let patience = int_opt opts "patience" (if quick then 3 else 5) in
+ let max_len = int_opt opts "max-len" (if quick then 6 else 7) in
+ let only = List.assoc_opt "only" opts in
+ let results =
+ Experiment.run_all
+ ~export_dir:(if render then Some (Filename.concat out_dir "steps") else None)
+ ~config ~beam_width:beam ~max_steps ~patience ~seed ~max_len ~only ()
+ in
+ if render then
+ List.iter
+ (fun (r : Experiment.result) ->
+ let base = Filename.concat out_dir r.Experiment.target.Experiment.name in
+ ignore
+ (Dot.render ~out_svg:(base ^ "-final.svg") ~out_dot:(base ^ "-final.dot")
+ (Dot.dot_grammar
+ ~title:(Printf.sprintf "%s (induced)" r.Experiment.target.Experiment.name)
+ r.Experiment.final_grammar)))
+ results;
+ Dot.write_file (Filename.concat out_dir "results.tsv")
+ (Experiment.to_tsv results);
+ Dot.write_file (Filename.concat out_dir "report.md")
+ (Experiment.report_markdown results);
+ List.iter
+ (fun (r : Experiment.result) ->
+ let base = Filename.concat out_dir r.Experiment.target.Experiment.name in
+ Dot.write_file (base ^ "-final.dot")
+ (Dot.dot_grammar
+ ~title:(Printf.sprintf "%s (induced)" r.Experiment.target.Experiment.name)
+ r.Experiment.final_grammar);
+ Dot.write_file (base ^ "-final.txt")
+ (Grammar.to_string r.Experiment.final_grammar))
+ results;
+ print_endline (Experiment.to_tsv results);
+ Printf.printf "\nwrote %s\n" out_dir
+
+let () =
+ match Array.to_list Sys.argv with
+ | _ :: cmd :: rest -> (
+ let opts = parse_args rest in
+ match cmd with
+ | "induce" -> cmd_induce opts
+ | "score" -> cmd_score opts
+ | "parse" -> cmd_parse opts
+ | "dot" -> cmd_dot opts
+ | "experiments" -> cmd_experiments opts
+ | "help" | "-h" | "--help" -> usage ()
+ | other ->
+ Printf.eprintf "unknown command: %s\n\n" other;
+ usage ();
+ exit 2)
+ | _ -> usage ()
diff --git a/dune-project b/dune-project
new file mode 100644
index 0000000..1a19819
--- /dev/null
+++ b/dune-project
@@ -0,0 +1,2 @@
+(lang dune 3.0)
+(name scfg_induction)
diff --git a/examples/anbn-paper.corpus b/examples/anbn-paper.corpus
new file mode 100644
index 0000000..4c3b14b
--- /dev/null
+++ b/examples/anbn-paper.corpus
@@ -0,0 +1,7 @@
+# The paper's running example {ab, aabb, aaabbb}, counts summing to 50.
+# Cook et al.'s exact relative frequencies are not reprinted in the paper, so a
+# decaying distribution is used; uniform counts (e.g. 17/17/16) correctly block
+# the inductive leap under the Bayesian criterion (see the report).
+30: a b
+15: a a b b
+5: a a a b b b
diff --git a/examples/palindrome-paper.grammar b/examples/palindrome-paper.grammar
new file mode 100644
index 0000000..a9f1691
--- /dev/null
+++ b/examples/palindrome-paper.grammar
@@ -0,0 +1,4 @@
+# Paper, Table 1, row "wcwR, w in {a,b}*": S -> c | aSa | bSb
+S -> c [0.333]
+S -> a S a [0.333]
+S -> b S b [0.333]
diff --git a/experiments/a2n-final.svg b/experiments/a2n-final.svg
new file mode 100644
index 0000000..7cd710f
--- /dev/null
+++ b/experiments/a2n-final.svg
@@ -0,0 +1,106 @@
+
+
+
+
+
diff --git a/experiments/a2n-final.txt b/experiments/a2n-final.txt
new file mode 100644
index 0000000..d7dfd86
--- /dev/null
+++ b/experiments/a2n-final.txt
@@ -0,0 +1,4 @@
+# start = S, 2 nonterminals, 1 terminals, 3 productions
+S -> S S [0.600]
+S -> T_a T_a [0.400]
+T_a -> a [1.000]
\ No newline at end of file
diff --git a/experiments/abn-final.svg b/experiments/abn-final.svg
new file mode 100644
index 0000000..36af491
--- /dev/null
+++ b/experiments/abn-final.svg
@@ -0,0 +1,137 @@
+
+
+
+
+
diff --git a/experiments/abn-final.txt b/experiments/abn-final.txt
new file mode 100644
index 0000000..060f4ba
--- /dev/null
+++ b/experiments/abn-final.txt
@@ -0,0 +1,5 @@
+# start = S, 3 nonterminals, 2 terminals, 4 productions
+S -> S S [0.600]
+T_b -> b [1.000]
+S -> T_a T_b [0.400]
+T_a -> a [1.000]
\ No newline at end of file
diff --git a/experiments/addition-final.svg b/experiments/addition-final.svg
new file mode 100644
index 0000000..47085c7
--- /dev/null
+++ b/experiments/addition-final.svg
@@ -0,0 +1,210 @@
+
+
+
+
+
diff --git a/experiments/addition-final.txt b/experiments/addition-final.txt
new file mode 100644
index 0000000..02c52ad
--- /dev/null
+++ b/experiments/addition-final.txt
@@ -0,0 +1,8 @@
+# start = S, 4 nonterminals, 5 terminals, 7 productions
+S -> b [0.143]
+T_) -> ) [1.000]
+S -> T_( S T_) [0.571]
+S -> S T_+ S [0.143]
+T_( -> ( [1.000]
+S -> a [0.143]
+T_+ -> + [1.000]
\ No newline at end of file
diff --git a/experiments/anbn-final.svg b/experiments/anbn-final.svg
new file mode 100644
index 0000000..f6f8574
--- /dev/null
+++ b/experiments/anbn-final.svg
@@ -0,0 +1,144 @@
+
+
+
+
+
diff --git a/experiments/anbn-final.txt b/experiments/anbn-final.txt
new file mode 100644
index 0000000..1763c20
--- /dev/null
+++ b/experiments/anbn-final.txt
@@ -0,0 +1,5 @@
+# start = S, 3 nonterminals, 2 terminals, 4 productions
+T_b -> b [1.000]
+S -> T_a S T_b [0.786]
+S -> T_a T_b [0.214]
+T_a -> a [1.000]
\ No newline at end of file
diff --git a/experiments/basic_english-final.svg b/experiments/basic_english-final.svg
new file mode 100644
index 0000000..4f2b5c2
--- /dev/null
+++ b/experiments/basic_english-final.svg
@@ -0,0 +1,690 @@
+
+
+
+
+
diff --git a/experiments/basic_english-final.txt b/experiments/basic_english-final.txt
new file mode 100644
index 0000000..64c023a
--- /dev/null
+++ b/experiments/basic_english-final.txt
@@ -0,0 +1,23 @@
+# start = S, 11 nonterminals, 17 terminals, 22 productions
+T_here -> here [0.500]
+S -> T_he T_is T_here [0.420]
+T_they -> they [0.333]
+S -> T_I T_am T_here [0.080]
+T_he -> she [0.333]
+T_that -> that [0.500]
+T_he -> it [0.333]
+T_am -> am [1.000]
+T_is -> is [1.000]
+T_man -> woman [0.500]
+S -> T_that T_man T_is T_here [0.080]
+T_that -> this [0.500]
+T_they -> we [0.333]
+T_are -> are [1.000]
+S -> T_that T_is T_a T_man [0.160]
+T_here -> there [0.500]
+T_he -> he [0.333]
+S -> T_they T_are T_here [0.260]
+T_a -> a [1.000]
+T_man -> man [0.500]
+T_I -> I [1.000]
+T_they -> you [0.333]
\ No newline at end of file
diff --git a/experiments/expr-final.svg b/experiments/expr-final.svg
new file mode 100644
index 0000000..060838a
--- /dev/null
+++ b/experiments/expr-final.svg
@@ -0,0 +1,230 @@
+
+
+
+
+
diff --git a/experiments/expr-final.txt b/experiments/expr-final.txt
new file mode 100644
index 0000000..66cade4
--- /dev/null
+++ b/experiments/expr-final.txt
@@ -0,0 +1,9 @@
+# start = S, 4 nonterminals, 6 terminals, 8 productions
+T_* -> * [0.500]
+S -> b [0.200]
+T_) -> ) [1.000]
+S -> T_( S T_) [0.200]
+T_* -> + [0.500]
+T_( -> ( [1.000]
+S -> a [0.200]
+S -> S T_* S [0.400]
\ No newline at end of file
diff --git a/experiments/nested_anbn-final.svg b/experiments/nested_anbn-final.svg
new file mode 100644
index 0000000..5ea88dd
--- /dev/null
+++ b/experiments/nested_anbn-final.svg
@@ -0,0 +1,142 @@
+
+
+
+
+
diff --git a/experiments/nested_anbn-final.txt b/experiments/nested_anbn-final.txt
new file mode 100644
index 0000000..fb5a9e8
--- /dev/null
+++ b/experiments/nested_anbn-final.txt
@@ -0,0 +1,5 @@
+# start = S, 3 nonterminals, 3 terminals, 4 productions
+T_b -> b [1.000]
+S -> c [0.059]
+S -> T_a S T_b [0.941]
+T_a -> a [1.000]
\ No newline at end of file
diff --git a/experiments/nested_parens-final.svg b/experiments/nested_parens-final.svg
new file mode 100644
index 0000000..23cd33f
--- /dev/null
+++ b/experiments/nested_parens-final.svg
@@ -0,0 +1,158 @@
+
+
+
+
+
diff --git a/experiments/nested_parens-final.txt b/experiments/nested_parens-final.txt
new file mode 100644
index 0000000..134830f
--- /dev/null
+++ b/experiments/nested_parens-final.txt
@@ -0,0 +1,6 @@
+# start = S, 3 nonterminals, 3 terminals, 5 productions
+S -> S S [0.300]
+T_) -> ) [1.000]
+S -> T_( S T_) [0.600]
+T_( -> ( [1.000]
+S -> a [0.100]
\ No newline at end of file
diff --git a/experiments/palindrome-final.svg b/experiments/palindrome-final.svg
new file mode 100644
index 0000000..cbfd555
--- /dev/null
+++ b/experiments/palindrome-final.svg
@@ -0,0 +1,176 @@
+
+
+
+
+
diff --git a/experiments/palindrome-final.txt b/experiments/palindrome-final.txt
new file mode 100644
index 0000000..7e8f6a6
--- /dev/null
+++ b/experiments/palindrome-final.txt
@@ -0,0 +1,6 @@
+# start = S, 3 nonterminals, 3 terminals, 5 productions
+T_b -> b [1.000]
+S -> c [0.083]
+S -> T_b S T_b [0.417]
+T_a -> a [1.000]
+S -> T_a S T_a [0.500]
\ No newline at end of file
diff --git a/experiments/parens-final.svg b/experiments/parens-final.svg
new file mode 100644
index 0000000..f746a59
--- /dev/null
+++ b/experiments/parens-final.svg
@@ -0,0 +1,160 @@
+
+
+
+
+
diff --git a/experiments/parens-final.txt b/experiments/parens-final.txt
new file mode 100644
index 0000000..3691e42
--- /dev/null
+++ b/experiments/parens-final.txt
@@ -0,0 +1,6 @@
+# start = S, 3 nonterminals, 2 terminals, 5 productions
+S -> S S [0.190]
+T_) -> ) [1.000]
+S -> T_( T_) [0.429]
+S -> T_( S T_) [0.381]
+T_( -> ( [1.000]
\ No newline at end of file
diff --git a/experiments/results.tsv b/experiments/results.tsv
new file mode 100644
index 0000000..57ba0af
--- /dev/null
+++ b/experiments/results.tsv
@@ -0,0 +1,12 @@
+group name nts prods initial_score final_score heldout_avg unparseable covers precision exact recursive steps iterations
+paper parens 3 5 -124.812 -100.91 -0.437804 0 true 1.0000 true true 2 7
+paper a2n 2 3 -60.667 -54.4321 -0.290242 0 true 1.0000 true true 2 7
+paper abn 3 4 -68.0175 -58.5388 -0.290242 0 true 1.0000 true true 2 7
+paper anbn 3 4 -76.1383 -68.7322 -0.336244 0 true 1.0000 true true 2 7
+paper palindrome 3 5 -198.878 -120.804 -0.716536 0 true 1.0000 true true 1 6
+paper addition 4 7 -203.571 -134.296 -1.11968 0 true 1.0000 true true 2 7
+paper shape 5 7 -251.225 -140.646 -0.646156 0 true 1.0000 true true 4 9
+paper basic_english 11 22 -470.394 -286.801 -0.986128 0 true 1.0000 true false 7 12
+synthetic nested_parens 3 5 -267.458 -149.867 -0.793712 0 true 1.0000 true true 1 6
+synthetic nested_anbn 3 4 -93.6716 -68.1568 -0.481173 0 true 1.0000 true true 1 6
+synthetic expr 4 8 -424.161 -223.342 -1.29086 0 true 1.0000 true true 3 8
\ No newline at end of file
diff --git a/experiments/shape-final.svg b/experiments/shape-final.svg
new file mode 100644
index 0000000..c6497ba
--- /dev/null
+++ b/experiments/shape-final.svg
@@ -0,0 +1,233 @@
+
+
+
+
+
diff --git a/experiments/shape-final.txt b/experiments/shape-final.txt
new file mode 100644
index 0000000..424228d
--- /dev/null
+++ b/experiments/shape-final.txt
@@ -0,0 +1,8 @@
+# start = S, 5 nonterminals, 4 terminals, 7 productions
+S -> T_d T_a [0.273]
+T_b -> b [1.000]
+T_a -> T_c T_a [0.917]
+T_c -> c [1.000]
+T_d -> d [1.000]
+T_a -> a [0.083]
+S -> T_b T_a S [0.727]
\ No newline at end of file
diff --git a/experiments/steps/a2n-step001-after.svg b/experiments/steps/a2n-step001-after.svg
new file mode 100644
index 0000000..cedcc93
--- /dev/null
+++ b/experiments/steps/a2n-step001-after.svg
@@ -0,0 +1,166 @@
+
+
+
+
+
diff --git a/experiments/steps/a2n-step001-before.svg b/experiments/steps/a2n-step001-before.svg
new file mode 100644
index 0000000..38310b1
--- /dev/null
+++ b/experiments/steps/a2n-step001-before.svg
@@ -0,0 +1,175 @@
+
+
+
+
+
diff --git a/experiments/steps/a2n-step002-after.svg b/experiments/steps/a2n-step002-after.svg
new file mode 100644
index 0000000..883d965
--- /dev/null
+++ b/experiments/steps/a2n-step002-after.svg
@@ -0,0 +1,106 @@
+
+
+
+
+
diff --git a/experiments/steps/a2n-step002-before.svg b/experiments/steps/a2n-step002-before.svg
new file mode 100644
index 0000000..30e1964
--- /dev/null
+++ b/experiments/steps/a2n-step002-before.svg
@@ -0,0 +1,166 @@
+
+
+
+
+
diff --git a/experiments/steps/abn-step001-after.svg b/experiments/steps/abn-step001-after.svg
new file mode 100644
index 0000000..04b4369
--- /dev/null
+++ b/experiments/steps/abn-step001-after.svg
@@ -0,0 +1,197 @@
+
+
+
+
+
diff --git a/experiments/steps/abn-step001-before.svg b/experiments/steps/abn-step001-before.svg
new file mode 100644
index 0000000..78d78bd
--- /dev/null
+++ b/experiments/steps/abn-step001-before.svg
@@ -0,0 +1,206 @@
+
+
+
+
+
diff --git a/experiments/steps/abn-step002-after.svg b/experiments/steps/abn-step002-after.svg
new file mode 100644
index 0000000..09922e6
--- /dev/null
+++ b/experiments/steps/abn-step002-after.svg
@@ -0,0 +1,137 @@
+
+
+
+
+
diff --git a/experiments/steps/abn-step002-before.svg b/experiments/steps/abn-step002-before.svg
new file mode 100644
index 0000000..86c4727
--- /dev/null
+++ b/experiments/steps/abn-step002-before.svg
@@ -0,0 +1,197 @@
+
+
+
+
+
diff --git a/experiments/steps/addition-step001-after.svg b/experiments/steps/addition-step001-after.svg
new file mode 100644
index 0000000..b6ea144
--- /dev/null
+++ b/experiments/steps/addition-step001-after.svg
@@ -0,0 +1,360 @@
+
+
+
+
+
diff --git a/experiments/steps/addition-step001-before.svg b/experiments/steps/addition-step001-before.svg
new file mode 100644
index 0000000..da7e347
--- /dev/null
+++ b/experiments/steps/addition-step001-before.svg
@@ -0,0 +1,592 @@
+
+
+
+
+
diff --git a/experiments/steps/addition-step002-after.svg b/experiments/steps/addition-step002-after.svg
new file mode 100644
index 0000000..851156e
--- /dev/null
+++ b/experiments/steps/addition-step002-after.svg
@@ -0,0 +1,210 @@
+
+
+
+
+
diff --git a/experiments/steps/addition-step002-before.svg b/experiments/steps/addition-step002-before.svg
new file mode 100644
index 0000000..f12fbd0
--- /dev/null
+++ b/experiments/steps/addition-step002-before.svg
@@ -0,0 +1,360 @@
+
+
+
+
+
diff --git a/experiments/steps/anbn-step001-after.svg b/experiments/steps/anbn-step001-after.svg
new file mode 100644
index 0000000..88bbfbe
--- /dev/null
+++ b/experiments/steps/anbn-step001-after.svg
@@ -0,0 +1,218 @@
+
+
+
+
+
diff --git a/experiments/steps/anbn-step001-before.svg b/experiments/steps/anbn-step001-before.svg
new file mode 100644
index 0000000..fa80939
--- /dev/null
+++ b/experiments/steps/anbn-step001-before.svg
@@ -0,0 +1,206 @@
+
+
+
+
+
diff --git a/experiments/steps/anbn-step002-after.svg b/experiments/steps/anbn-step002-after.svg
new file mode 100644
index 0000000..18c88d3
--- /dev/null
+++ b/experiments/steps/anbn-step002-after.svg
@@ -0,0 +1,144 @@
+
+
+
+
+
diff --git a/experiments/steps/anbn-step002-before.svg b/experiments/steps/anbn-step002-before.svg
new file mode 100644
index 0000000..ecab941
--- /dev/null
+++ b/experiments/steps/anbn-step002-before.svg
@@ -0,0 +1,218 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step001-after.svg b/experiments/steps/basic_english-step001-after.svg
new file mode 100644
index 0000000..06b2b7d
--- /dev/null
+++ b/experiments/steps/basic_english-step001-after.svg
@@ -0,0 +1,1026 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step001-before.svg b/experiments/steps/basic_english-step001-before.svg
new file mode 100644
index 0000000..cbc1d2b
--- /dev/null
+++ b/experiments/steps/basic_english-step001-before.svg
@@ -0,0 +1,1311 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step002-after.svg b/experiments/steps/basic_english-step002-after.svg
new file mode 100644
index 0000000..1de5fbb
--- /dev/null
+++ b/experiments/steps/basic_english-step002-after.svg
@@ -0,0 +1,938 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step002-before.svg b/experiments/steps/basic_english-step002-before.svg
new file mode 100644
index 0000000..e054022
--- /dev/null
+++ b/experiments/steps/basic_english-step002-before.svg
@@ -0,0 +1,1026 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step003-after.svg b/experiments/steps/basic_english-step003-after.svg
new file mode 100644
index 0000000..2d37c96
--- /dev/null
+++ b/experiments/steps/basic_english-step003-after.svg
@@ -0,0 +1,850 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step003-before.svg b/experiments/steps/basic_english-step003-before.svg
new file mode 100644
index 0000000..b8c03d2
--- /dev/null
+++ b/experiments/steps/basic_english-step003-before.svg
@@ -0,0 +1,938 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step004-after.svg b/experiments/steps/basic_english-step004-after.svg
new file mode 100644
index 0000000..98905cb
--- /dev/null
+++ b/experiments/steps/basic_english-step004-after.svg
@@ -0,0 +1,810 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step004-before.svg b/experiments/steps/basic_english-step004-before.svg
new file mode 100644
index 0000000..4527580
--- /dev/null
+++ b/experiments/steps/basic_english-step004-before.svg
@@ -0,0 +1,850 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step005-after.svg b/experiments/steps/basic_english-step005-after.svg
new file mode 100644
index 0000000..7b79ff7
--- /dev/null
+++ b/experiments/steps/basic_english-step005-after.svg
@@ -0,0 +1,770 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step005-before.svg b/experiments/steps/basic_english-step005-before.svg
new file mode 100644
index 0000000..65c8a5c
--- /dev/null
+++ b/experiments/steps/basic_english-step005-before.svg
@@ -0,0 +1,810 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step006-after.svg b/experiments/steps/basic_english-step006-after.svg
new file mode 100644
index 0000000..c8ceb18
--- /dev/null
+++ b/experiments/steps/basic_english-step006-after.svg
@@ -0,0 +1,730 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step006-before.svg b/experiments/steps/basic_english-step006-before.svg
new file mode 100644
index 0000000..66f601e
--- /dev/null
+++ b/experiments/steps/basic_english-step006-before.svg
@@ -0,0 +1,770 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step007-after.svg b/experiments/steps/basic_english-step007-after.svg
new file mode 100644
index 0000000..0ec4f4c
--- /dev/null
+++ b/experiments/steps/basic_english-step007-after.svg
@@ -0,0 +1,690 @@
+
+
+
+
+
diff --git a/experiments/steps/basic_english-step007-before.svg b/experiments/steps/basic_english-step007-before.svg
new file mode 100644
index 0000000..31b43f2
--- /dev/null
+++ b/experiments/steps/basic_english-step007-before.svg
@@ -0,0 +1,730 @@
+
+
+
+
+
diff --git a/experiments/steps/expr-step001-after.svg b/experiments/steps/expr-step001-after.svg
new file mode 100644
index 0000000..a6db81f
--- /dev/null
+++ b/experiments/steps/expr-step001-after.svg
@@ -0,0 +1,290 @@
+
+
+
+
+
diff --git a/experiments/steps/expr-step001-before.svg b/experiments/steps/expr-step001-before.svg
new file mode 100644
index 0000000..4aaf8a6
--- /dev/null
+++ b/experiments/steps/expr-step001-before.svg
@@ -0,0 +1,1214 @@
+
+
+
+
+
diff --git a/experiments/steps/expr-step002-after.svg b/experiments/steps/expr-step002-after.svg
new file mode 100644
index 0000000..7e32547
--- /dev/null
+++ b/experiments/steps/expr-step002-after.svg
@@ -0,0 +1,256 @@
+
+
+
+
+
diff --git a/experiments/steps/expr-step002-before.svg b/experiments/steps/expr-step002-before.svg
new file mode 100644
index 0000000..87b8afb
--- /dev/null
+++ b/experiments/steps/expr-step002-before.svg
@@ -0,0 +1,290 @@
+
+
+
+
+
diff --git a/experiments/steps/expr-step003-after.svg b/experiments/steps/expr-step003-after.svg
new file mode 100644
index 0000000..0516a2d
--- /dev/null
+++ b/experiments/steps/expr-step003-after.svg
@@ -0,0 +1,230 @@
+
+
+
+
+
diff --git a/experiments/steps/expr-step003-before.svg b/experiments/steps/expr-step003-before.svg
new file mode 100644
index 0000000..6b245cd
--- /dev/null
+++ b/experiments/steps/expr-step003-before.svg
@@ -0,0 +1,256 @@
+
+
+
+
+
diff --git a/experiments/steps/nested_anbn-step001-after.svg b/experiments/steps/nested_anbn-step001-after.svg
new file mode 100644
index 0000000..10f3360
--- /dev/null
+++ b/experiments/steps/nested_anbn-step001-after.svg
@@ -0,0 +1,142 @@
+
+
+
+
+
diff --git a/experiments/steps/nested_anbn-step001-before.svg b/experiments/steps/nested_anbn-step001-before.svg
new file mode 100644
index 0000000..1d2d390
--- /dev/null
+++ b/experiments/steps/nested_anbn-step001-before.svg
@@ -0,0 +1,278 @@
+
+
+
+
+
diff --git a/experiments/steps/nested_parens-step001-after.svg b/experiments/steps/nested_parens-step001-after.svg
new file mode 100644
index 0000000..7779dba
--- /dev/null
+++ b/experiments/steps/nested_parens-step001-after.svg
@@ -0,0 +1,158 @@
+
+
+
+
+
diff --git a/experiments/steps/nested_parens-step001-before.svg b/experiments/steps/nested_parens-step001-before.svg
new file mode 100644
index 0000000..ef4b623
--- /dev/null
+++ b/experiments/steps/nested_parens-step001-before.svg
@@ -0,0 +1,829 @@
+
+
+
+
+
diff --git a/experiments/steps/palindrome-step001-after.svg b/experiments/steps/palindrome-step001-after.svg
new file mode 100644
index 0000000..62c5609
--- /dev/null
+++ b/experiments/steps/palindrome-step001-after.svg
@@ -0,0 +1,176 @@
+
+
+
+
+
diff --git a/experiments/steps/palindrome-step001-before.svg b/experiments/steps/palindrome-step001-before.svg
new file mode 100644
index 0000000..7ad1702
--- /dev/null
+++ b/experiments/steps/palindrome-step001-before.svg
@@ -0,0 +1,580 @@
+
+
+
+
+
diff --git a/experiments/steps/parens-step001-after.svg b/experiments/steps/parens-step001-after.svg
new file mode 100644
index 0000000..cc7cbc0
--- /dev/null
+++ b/experiments/steps/parens-step001-after.svg
@@ -0,0 +1,316 @@
+
+
+
+
+
diff --git a/experiments/steps/parens-step001-before.svg b/experiments/steps/parens-step001-before.svg
new file mode 100644
index 0000000..f8d5fea
--- /dev/null
+++ b/experiments/steps/parens-step001-before.svg
@@ -0,0 +1,346 @@
+
+
+
+
+
diff --git a/experiments/steps/parens-step002-after.svg b/experiments/steps/parens-step002-after.svg
new file mode 100644
index 0000000..4e07bec
--- /dev/null
+++ b/experiments/steps/parens-step002-after.svg
@@ -0,0 +1,160 @@
+
+
+
+
+
diff --git a/experiments/steps/parens-step002-before.svg b/experiments/steps/parens-step002-before.svg
new file mode 100644
index 0000000..ef7fc4e
--- /dev/null
+++ b/experiments/steps/parens-step002-before.svg
@@ -0,0 +1,316 @@
+
+
+
+
+
diff --git a/experiments/steps/shape-step001-after.svg b/experiments/steps/shape-step001-after.svg
new file mode 100644
index 0000000..8d97e39
--- /dev/null
+++ b/experiments/steps/shape-step001-after.svg
@@ -0,0 +1,698 @@
+
+
+
+
+
diff --git a/experiments/steps/shape-step001-before.svg b/experiments/steps/shape-step001-before.svg
new file mode 100644
index 0000000..148d01f
--- /dev/null
+++ b/experiments/steps/shape-step001-before.svg
@@ -0,0 +1,742 @@
+
+
+
+
+
diff --git a/experiments/steps/shape-step002-after.svg b/experiments/steps/shape-step002-after.svg
new file mode 100644
index 0000000..354c94c
--- /dev/null
+++ b/experiments/steps/shape-step002-after.svg
@@ -0,0 +1,295 @@
+
+
+
+
+
diff --git a/experiments/steps/shape-step002-before.svg b/experiments/steps/shape-step002-before.svg
new file mode 100644
index 0000000..6f6858d
--- /dev/null
+++ b/experiments/steps/shape-step002-before.svg
@@ -0,0 +1,698 @@
+
+
+
+
+
diff --git a/experiments/steps/shape-step003-after.svg b/experiments/steps/shape-step003-after.svg
new file mode 100644
index 0000000..34ffd4a
--- /dev/null
+++ b/experiments/steps/shape-step003-after.svg
@@ -0,0 +1,307 @@
+
+
+
+
+
diff --git a/experiments/steps/shape-step003-before.svg b/experiments/steps/shape-step003-before.svg
new file mode 100644
index 0000000..555d77b
--- /dev/null
+++ b/experiments/steps/shape-step003-before.svg
@@ -0,0 +1,295 @@
+
+
+
+
+
diff --git a/experiments/steps/shape-step004-after.svg b/experiments/steps/shape-step004-after.svg
new file mode 100644
index 0000000..e5457a9
--- /dev/null
+++ b/experiments/steps/shape-step004-after.svg
@@ -0,0 +1,233 @@
+
+
+
+
+
diff --git a/experiments/steps/shape-step004-before.svg b/experiments/steps/shape-step004-before.svg
new file mode 100644
index 0000000..50286b6
--- /dev/null
+++ b/experiments/steps/shape-step004-before.svg
@@ -0,0 +1,307 @@
+
+
+
+
+
diff --git a/lib/corpus.ml b/lib/corpus.ml
new file mode 100644
index 0000000..aca4e73
--- /dev/null
+++ b/lib/corpus.ml
@@ -0,0 +1,67 @@
+type sample = { tokens : string list; count : float }
+
+type t = sample list
+
+let total_count corpus =
+ List.fold_left (fun s x -> s +. x.count) 0.0 corpus
+
+let distinct corpus = List.length corpus
+
+let tokens_to_string toks = String.concat " " toks
+
+let of_string text =
+ let lines = String.split_on_char '\n' text in
+ let acc = Hashtbl.create 64 in
+ let order = ref [] in
+ List.iter
+ (fun line ->
+ let line = String.trim line in
+ if String.length line > 0 && line.[0] <> '#' then begin
+
+ let count, rest =
+ match String.index_opt line ':' with
+ | Some i -> (
+ match
+ float_of_string_opt (String.trim (String.sub line 0 i))
+ with
+ | Some c ->
+ ( c,
+ String.sub line (i + 1) (String.length line - i - 1) )
+ | None -> (1.0, line))
+ | None -> (1.0, line)
+ in
+ let toks =
+ String.split_on_char ' ' (String.trim rest)
+ |> List.filter (fun s -> String.length s > 0)
+ in
+ if toks <> [] then begin
+ let key = tokens_to_string toks in
+ (match Hashtbl.find_opt acc key with
+ | Some c -> Hashtbl.replace acc key (c +. count)
+ | None ->
+ Hashtbl.replace acc key count;
+ order := key :: !order);
+ ()
+ end
+ end)
+ lines;
+ List.rev !order
+ |> List.map (fun key ->
+ let toks =
+ String.split_on_char ' ' key
+ |> List.filter (fun s -> String.length s > 0)
+ in
+ { tokens = toks; count = Hashtbl.find acc key })
+
+let to_string corpus =
+ List.map
+ (fun s -> Printf.sprintf "%g: %s" s.count (tokens_to_string s.tokens))
+ corpus
+ |> String.concat "\n"
+
+let of_file path =
+ let ic = open_in path in
+ let n = in_channel_length ic in
+ let s = really_input_string ic n in
+ close_in ic;
+ of_string s
diff --git a/lib/dot.ml b/lib/dot.ml
new file mode 100644
index 0000000..0b2c71f
--- /dev/null
+++ b/lib/dot.ml
@@ -0,0 +1,143 @@
+let escape s =
+ let b = Buffer.create (String.length s) in
+ String.iter
+ (fun c ->
+ match c with
+ | '"' -> Buffer.add_string b "\\\""
+ | '\\' -> Buffer.add_string b "\\\\"
+ | '\n' -> Buffer.add_string b "\\n"
+ | c -> Buffer.add_char b c)
+ s;
+ Buffer.contents b
+
+let sanitize s =
+ String.map
+ (fun c ->
+ if
+ (c >= 'a' && c <= 'z')
+ || (c >= 'A' && c <= 'Z')
+ || (c >= '0' && c <= '9')
+ || c = '_'
+ then c
+ else '_')
+ s
+
+let nt_id name = "nt_" ^ sanitize name
+
+let prob_of probs lhs rhs =
+ match Hashtbl.find_opt probs lhs with
+ | None -> 0.0
+ | Some l -> (
+ match List.find_opt (fun (q, _) -> Grammar.equal_rhs q.Grammar.rhs rhs) l with
+ | Some (_, x) -> x
+ | None -> 0.0)
+
+let prod_key p =
+ (p.Grammar.lhs, String.concat " " (List.map Grammar.symbol_to_string p.Grammar.rhs))
+
+let dot_grammar ?(highlight_nts = []) ?(highlight_prods = [])
+ ?(title = "grammar") g =
+ let probs = Grammar.probabilities g in
+ let hl_nt name = List.mem name highlight_nts in
+ let hl_prods = List.map prod_key highlight_prods in
+ let buf = Buffer.create 2048 in
+ Buffer.add_string buf "digraph grammar {\n";
+ Buffer.add_string buf " rankdir=LR;\n";
+ Buffer.add_string buf " labelloc=\"t\";\n";
+ Buffer.add_string buf (Printf.sprintf " label=\"%s\";\n" (escape title));
+ Buffer.add_string buf
+ " node [fontname=\"Helvetica\", fontsize=10];\n edge [fontname=\"Helvetica\", fontsize=9];\n";
+ List.iter
+ (fun nt ->
+ let is_start = String.equal nt g.Grammar.start in
+ let hl = hl_nt nt in
+ Buffer.add_string buf
+ (Printf.sprintf
+ " %s [label=\"%s\", shape=%s%s%s];\n"
+ (nt_id nt) (escape nt)
+ (if is_start then "doublecircle" else "ellipse")
+ (if hl then ", color=red, penwidth=2.5, fontcolor=red" else "")
+ (if is_start then ", style=bold" else "")))
+ (Grammar.nonterminals g);
+ List.iter
+ (fun t ->
+ Buffer.add_string buf
+ (Printf.sprintf
+ " tm_%s [label=\"%s\", shape=plaintext, fontcolor=\"#666666\"];\n"
+ (sanitize t) (escape t)))
+ (Grammar.terminals g);
+ List.iteri
+ (fun i p ->
+ let pr = prob_of probs p.Grammar.lhs p.Grammar.rhs in
+ let hl = List.mem (prod_key p) hl_prods in
+ let rhs_text =
+ String.concat " " (List.map Grammar.symbol_to_string p.Grammar.rhs)
+ in
+ Buffer.add_string buf
+ (Printf.sprintf
+ " pr_%d [label=\"%s\\np=%.3f\", shape=box%s];\n"
+ i (escape rhs_text) pr
+ (if hl then ", color=red, penwidth=2.5, fontcolor=red" else ""));
+ Buffer.add_string buf
+ (Printf.sprintf " %s -> pr_%d;\n" (nt_id p.Grammar.lhs) i);
+ List.iteri
+ (fun k s ->
+ let target =
+ match s with
+ | Grammar.Nonterm x -> nt_id x
+ | Grammar.Term t -> "tm_" ^ sanitize t
+ in
+ Buffer.add_string buf
+ (Printf.sprintf " pr_%d -> %s [label=\"%d\"];\n" i target (k + 1)))
+ p.Grammar.rhs)
+ g.Grammar.productions;
+ Buffer.add_string buf "}\n";
+ Buffer.contents buf
+
+let dot_tree ?(title = "parse") tree =
+ let buf = Buffer.create 1024 in
+ Buffer.add_string buf "digraph parse {\n";
+ Buffer.add_string buf " rankdir=TB;\n";
+ Buffer.add_string buf (Printf.sprintf " label=\"%s\";\n" (escape title));
+ Buffer.add_string buf
+ " node [fontname=\"Helvetica\", fontsize=10, shape=ellipse];\n";
+ let counter = ref 0 in
+ let rec go tree =
+ let id = !counter in
+ incr counter;
+ match tree with
+ | Parse.Leaf s ->
+ Buffer.add_string buf
+ (Printf.sprintf " n%d [label=\"%s\", shape=plaintext, fontcolor=\"#666666\"];\n"
+ id (escape s));
+ id
+ | Parse.Node (nt, children) ->
+ Buffer.add_string buf
+ (Printf.sprintf " n%d [label=\"%s\"];\n" id (escape nt));
+ List.iteri
+ (fun k child ->
+ let cid = go child in
+ Buffer.add_string buf
+ (Printf.sprintf " n%d -> n%d [label=\"%d\"];\n" id cid (k + 1)))
+ children;
+ id
+ in
+ ignore (go tree);
+ Buffer.add_string buf "}\n";
+ Buffer.contents buf
+
+let write_file path contents =
+ let oc = open_out path in
+ output_string oc contents;
+ close_out oc
+
+let render ?out_svg ~out_dot dot_string =
+ write_file out_dot dot_string;
+ match out_svg with
+ | None -> None
+ | Some svg ->
+ let cmd =
+ Printf.sprintf "dot -Tsvg %s -o %s 2>/dev/null"
+ (Filename.quote out_dot) (Filename.quote svg)
+ in
+ if Sys.command cmd = 0 then Some svg else None
diff --git a/lib/dune b/lib/dune
new file mode 100644
index 0000000..8136f94
--- /dev/null
+++ b/lib/dune
@@ -0,0 +1,4 @@
+(library
+ (name scfg)
+ (flags (:standard -warn-error -a))
+ (modules grammar corpus parse scoring transform search dot experiment))
diff --git a/lib/experiment.ml b/lib/experiment.ml
new file mode 100644
index 0000000..81b1990
--- /dev/null
+++ b/lib/experiment.ml
@@ -0,0 +1,282 @@
+type target = {
+ name : string;
+ group : string;
+ desc : string;
+ grammar : Grammar.t;
+}
+
+let tm x = Grammar.Term x
+let nt x = Grammar.Nonterm x
+
+let tg group name desc start prods =
+ let productions =
+ List.map
+ (fun (lhs, rhs, count) -> { Grammar.lhs; rhs; count })
+ prods
+ in
+ { group; name; desc; grammar = Grammar.make ~start productions }
+
+let targets =
+ [
+ tg "paper" "parens" "balanced parentheses S -> () | (S) | SS" "S"
+ [ ("S", [ tm "("; tm ")" ], 1.0); ("S", [ tm "("; nt "S"; tm ")" ], 1.0);
+ ("S", [ nt "S"; nt "S" ], 1.0) ];
+ tg "paper" "a2n" "even-length strings of a S -> aa | SS" "S"
+ [ ("S", [ tm "a"; tm "a" ], 1.0); ("S", [ nt "S"; nt "S" ], 1.0) ];
+ tg "paper" "abn" "(ab)^n S -> ab | SS" "S"
+ [ ("S", [ tm "a"; tm "b" ], 1.0); ("S", [ nt "S"; nt "S" ], 1.0) ];
+ tg "paper" "anbn" "a^n b^n S -> ab | aSb" "S"
+ [ ("S", [ tm "a"; tm "b" ], 1.0); ("S", [ tm "a"; nt "S"; tm "b" ], 1.0) ];
+ tg "paper" "palindrome" "w c w^R S -> c | aSa | bSb" "S"
+ [ ("S", [ tm "c" ], 1.0); ("S", [ tm "a"; nt "S"; tm "a" ], 1.0);
+ ("S", [ tm "b"; nt "S"; tm "b" ], 1.0) ];
+ tg "paper" "addition" "addition strings S -> a | b | (S) | S+S" "S"
+ [ ("S", [ tm "a" ], 1.0); ("S", [ tm "b" ], 1.0);
+ ("S", [ tm "("; nt "S"; tm ")" ], 1.0);
+ ("S", [ nt "S"; tm "+"; nt "S" ], 1.0) ];
+ tg "paper" "shape" "shape grammar S -> dY | bYS ; Y -> a | cY" "S"
+ [ ("S", [ tm "d"; nt "Y" ], 1.0);
+ ("S", [ tm "b"; nt "Y"; nt "S" ], 1.0); ("Y", [ tm "a" ], 1.0);
+ ("Y", [ tm "c"; nt "Y" ], 1.0) ];
+ tg "paper" "basic_english"
+ "basic English S -> I am A | he T | ..." "S"
+ [ ("S", [ tm "I"; tm "am"; nt "A" ], 1.0);
+ ("S", [ tm "he"; nt "T" ], 1.0);
+ ("S", [ tm "she"; nt "T" ], 1.0);
+ ("S", [ tm "it"; nt "T" ], 1.0);
+ ("S", [ tm "they"; nt "V" ], 1.0);
+ ("S", [ tm "you"; nt "V" ], 1.0);
+ ("S", [ tm "we"; nt "V" ], 1.0);
+ ("S", [ tm "this"; nt "C" ], 1.0);
+ ("S", [ tm "that"; nt "C" ], 1.0);
+ ("T", [ tm "is"; nt "A" ], 1.0);
+ ("V", [ tm "are"; nt "A" ], 1.0);
+ ("A", [ tm "there" ], 1.0); ("A", [ tm "here" ], 1.0);
+ ("C", [ tm "is"; tm "a"; nt "Z" ], 1.0);
+ ("C", [ nt "Z"; nt "T" ], 1.0);
+ ("Z", [ tm "man" ], 1.0); ("Z", [ tm "woman" ], 1.0) ];
+ tg "synthetic" "nested_parens" "nested + concatenation S -> a | (S) | SS" "S"
+ [ ("S", [ tm "a" ], 1.0); ("S", [ tm "("; nt "S"; tm ")" ], 1.0);
+ ("S", [ nt "S"; nt "S" ], 1.0) ];
+ tg "synthetic" "nested_anbn" "nested a^n b^n S -> aSb | c" "S"
+ [ ("S", [ tm "a"; nt "S"; tm "b" ], 1.0); ("S", [ tm "c" ], 1.0) ];
+ tg "synthetic" "expr" "arithmetic S -> S+T | T ; T -> T*F | F ; F -> (S) | a | b" "S"
+ [ ("S", [ nt "S"; tm "+"; nt "T" ], 1.0); ("S", [ nt "T" ], 1.0);
+ ("T", [ nt "T"; tm "*"; nt "F" ], 1.0); ("T", [ nt "F" ], 1.0);
+ ("F", [ tm "("; nt "S"; tm ")" ], 1.0); ("F", [ tm "a" ], 1.0);
+ ("F", [ tm "b" ], 1.0) ];
+ ]
+
+let sample_one probs start rng max_pending =
+ let rec derive pending =
+ if List.length pending > max_pending then None
+ else
+ match pending with
+ | [] -> Some []
+ | Grammar.Term a :: rest -> (
+ match derive rest with Some l -> Some (a :: l) | None -> None)
+ | Grammar.Nonterm x :: rest -> (
+ match Hashtbl.find_opt probs x with
+ | None | Some [] -> None
+ | Some ps ->
+ let tot = List.fold_left (fun s (_, p) -> s +. p) 0.0 ps in
+ let r = Random.State.float rng tot in
+ let rec pick acc = function
+ | [] -> fst (List.hd ps)
+ | (p, pr) :: tl -> if acc +. pr >= r then p else pick (acc +. pr) tl
+ in
+ let p = pick 0.0 ps in
+ derive (p.Grammar.rhs @ rest))
+ in
+ derive [ Grammar.Nonterm start ]
+
+let sample_counts target draws seed max_len =
+ let probs = Grammar.probabilities target.grammar in
+ let rng = Random.State.make [| seed |] in
+ let tbl = Hashtbl.create 256 in
+ let order = ref [] in
+ let rec go i tries =
+ if i >= draws || tries > draws * 100 then ()
+ else
+ match sample_one probs target.grammar.Grammar.start rng (2 * max_len) with
+ | Some s when List.length s >= 1 && List.length s <= max_len ->
+ (match Hashtbl.find_opt tbl s with
+ | Some c -> Hashtbl.replace tbl s (c + 1)
+ | None ->
+ Hashtbl.replace tbl s 1;
+ order := s :: !order);
+ go (i + 1) (tries + 1)
+ | _ -> go i (tries + 1)
+ in
+ go 0 0;
+ List.rev !order
+ |> List.map (fun s ->
+ { Corpus.tokens = s; count = float_of_int (Hashtbl.find tbl s) })
+
+let initial_of_corpus corpus =
+ Grammar.initial_grammar
+ (List.map (fun s -> (s.Corpus.tokens, s.Corpus.count)) corpus)
+
+type result = {
+ target : target;
+ train : Corpus.t;
+ test : Corpus.t;
+ initial : Grammar.t;
+ final_grammar : Grammar.t;
+ initial_score : float;
+ final_score : float;
+ nts : int;
+ prods : int;
+ heldout_avg : float;
+ unparseable_test : int;
+ covers_target : bool;
+ precision : float;
+ exact_language : bool;
+ recursive : bool;
+ steps : int;
+ iterations : int;
+}
+
+let language_sets target final_grammar max_len =
+ let tlang = Grammar.language_up_to target.grammar max_len in
+ let flang = Grammar.language_up_to final_grammar max_len in
+ let covers = List.for_all (fun s -> List.mem s flang) tlang in
+ let precision =
+ if flang = [] then 0.0
+ else
+ float_of_int
+ (List.length (List.filter (fun s -> List.mem s tlang) flang))
+ /. float_of_int (List.length flang)
+ in
+ let exact =
+ List.length tlang = List.length flang
+ && List.for_all (fun s -> List.mem s flang) tlang
+ in
+ (covers, precision, exact)
+
+let run_one ?(export_dir = None) ~config ~beam_width ~max_steps ~patience ~seed
+ ~max_len target =
+ let train = sample_counts target 50 seed max_len in
+ let test = sample_counts target 200 (seed + 1000) max_len in
+ let initial = initial_of_corpus train in
+ let search =
+ Search.run ~config ~beam_width ~max_steps ~patience initial train
+ in
+ (match export_dir with
+ | None -> ()
+ | Some dir ->
+ List.iteri
+ (fun i (s : Search.step) ->
+ let ch =
+ Transform.change_of s.Search.parent s.Search.transform s.Search.result
+ in
+ let base =
+ Filename.concat dir
+ (Printf.sprintf "%s-step%03d" target.name (i + 1))
+ in
+ ignore
+ (Dot.render ~out_svg:(base ^ "-before.svg")
+ ~out_dot:(base ^ "-before.dot")
+ (Dot.dot_grammar ~highlight_nts:ch.Transform.before_nts
+ ~highlight_prods:ch.Transform.before_prods
+ ~title:
+ (Printf.sprintf "%s step %d before: %s" target.name (i + 1)
+ (Transform.describe s.Search.transform))
+ s.Search.parent));
+ ignore
+ (Dot.render ~out_svg:(base ^ "-after.svg")
+ ~out_dot:(base ^ "-after.dot")
+ (Dot.dot_grammar ~highlight_nts:ch.Transform.after_nts
+ ~highlight_prods:ch.Transform.after_prods
+ ~title:
+ (Printf.sprintf "%s step %d after: %s" target.name (i + 1)
+ (Transform.describe s.Search.transform))
+ s.Search.result)))
+ search.Search.steps);
+
+ let final_grammar = search.Search.best in
+ let heldout_avg, _, _, unparseable =
+ Scoring.heldout ~config final_grammar ~train test
+ in
+ let covers, precision, exact = language_sets target final_grammar max_len in
+ {
+ target;
+ train;
+ test;
+ initial;
+ final_grammar;
+ initial_score = search.Search.initial_score;
+ final_score = search.Search.best_score;
+ nts = List.length (Grammar.nonterminals final_grammar);
+ prods = List.length final_grammar.Grammar.productions;
+ heldout_avg;
+ unparseable_test = unparseable;
+ covers_target = covers;
+ precision;
+ exact_language = exact;
+ recursive = Grammar.is_recursive final_grammar;
+ steps = List.length search.Search.steps;
+ iterations = search.Search.iterations;
+ }
+
+let header () =
+ "group\tname\tnts\tprods\tinitial_score\tfinal_score\theldout_avg\tunparseable\t\
+ covers\tprecision\texact\trecursive\tsteps\titerations"
+
+let row r =
+ Printf.sprintf
+ "%s\t%s\t%d\t%d\t%.6g\t%.6g\t%.6g\t%d\t%b\t%.4f\t%b\t%b\t%d\t%d"
+ r.target.group r.target.name r.nts r.prods r.initial_score r.final_score
+ r.heldout_avg r.unparseable_test r.covers_target r.precision
+ r.exact_language r.recursive r.steps r.iterations
+
+let to_tsv results = String.concat "\n" (header () :: List.map row results)
+
+let report_markdown results =
+ let b = Buffer.create 4096 in
+ Buffer.add_string b "# SCFG induction experiments\n\n";
+ Buffer.add_string b
+ "Scores are natural-log posterior values. `heldout_avg` is the average \
+ log-likelihood per token on a held-out sample drawn from the same target \
+ grammar (higher is better). `precision` is the fraction of strings the \
+ induced grammar generates up to the comparison length that belong to the \
+ target language.\n\n";
+ Buffer.add_string b
+ "| group | target | nonterms | prods | initial | final | heldout/token | \
+ covers | precision | exact | recursive | steps |\n";
+ Buffer.add_string b
+ "|---|---|---|---|---|---|---|---|---|---|---|---|\n";
+ List.iter
+ (fun r ->
+ Buffer.add_string b
+ (Printf.sprintf
+ "| %s | %s | %d | %d | %.2f | %.2f | %.3f | %b | %.3f | %b | %b | %d |\n"
+ r.target.group r.target.name r.nts r.prods r.initial_score
+ r.final_score r.heldout_avg r.covers_target r.precision
+ r.exact_language r.recursive r.steps))
+ results;
+ Buffer.add_string b "\n## Final grammars\n\n";
+ List.iter
+ (fun r ->
+ Buffer.add_string b
+ (Printf.sprintf "### %s (%s)\n\n%s\n\n" r.target.name r.target.group
+ r.target.desc);
+ Buffer.add_string b "```\n";
+ Buffer.add_string b (Grammar.to_string r.final_grammar);
+ Buffer.add_string b "\n```\n\n")
+ results;
+ Buffer.contents b
+
+let run_all ?(export_dir = None) ?(config = Scoring.default_config)
+ ?(beam_width = 3) ?(max_steps = 25) ?(patience = 5) ?(seed = 1)
+ ?(max_len = 8) ?(only = None) () =
+ let ts =
+ match only with
+ | None -> targets
+ | Some name -> List.filter (fun t -> String.equal t.name name) targets
+ in
+ List.map
+ (fun target ->
+ run_one ~export_dir ~config ~beam_width ~max_steps ~patience ~seed ~max_len
+ target)
+ ts
diff --git a/lib/grammar.ml b/lib/grammar.ml
new file mode 100644
index 0000000..51b9433
--- /dev/null
+++ b/lib/grammar.ml
@@ -0,0 +1,522 @@
+type symbol = Term of string | Nonterm of string
+
+type production = {
+ lhs : string;
+ rhs : symbol list;
+ count : float;
+}
+
+type t = {
+ start : string;
+ productions : production list;
+}
+
+module SSet = Set.Make (String)
+
+let symbol_to_string = function Term s -> s | Nonterm s -> s
+
+let symbol_key = function Term s -> "t:" ^ s | Nonterm s -> "n:" ^ s
+
+let symbol_is_nonterm = function Nonterm _ -> true | Term _ -> false
+
+let equal_symbol a b =
+ match (a, b) with
+ | Term x, Term y | Nonterm x, Nonterm y -> String.equal x y
+ | _ -> false
+
+let equal_rhs a b =
+ List.length a = List.length b && List.for_all2 equal_symbol a b
+
+let rhs_to_string rhs = String.concat " " (List.map symbol_to_string rhs)
+
+let production_to_string p =
+ Printf.sprintf "%s -> %s" p.lhs (rhs_to_string p.rhs)
+
+let make ~start productions = { start; productions }
+
+let nonterminals g =
+ let acc = ref (SSet.singleton g.start) in
+ List.iter
+ (fun p ->
+ acc := SSet.add p.lhs !acc;
+ List.iter
+ (function Nonterm x -> acc := SSet.add x !acc | Term _ -> ())
+ p.rhs)
+ g.productions;
+ SSet.elements !acc
+
+let terminals g =
+ let acc = ref SSet.empty in
+ List.iter
+ (fun p ->
+ List.iter
+ (function Term x -> acc := SSet.add x !acc | Nonterm _ -> ())
+ p.rhs)
+ g.productions;
+ SSet.elements !acc
+
+let productions_of g lhs =
+ List.filter (fun p -> String.equal p.lhs lhs) g.productions
+
+let total_count prods = List.fold_left (fun s p -> s +. p.count) 0.0 prods
+
+let probabilities g =
+ let by_lhs = Hashtbl.create 64 in
+ List.iter
+ (fun p ->
+ let l =
+ match Hashtbl.find_opt by_lhs p.lhs with Some l -> l | None -> []
+ in
+ Hashtbl.replace by_lhs p.lhs (p :: l))
+ g.productions;
+ let tbl = Hashtbl.create 64 in
+ Hashtbl.iter
+ (fun lhs ps ->
+ let ps = List.rev ps in
+ let tot = total_count ps in
+ let k = List.length ps in
+ let probs =
+ if tot > 0.0 then List.map (fun p -> (p, p.count /. tot)) ps
+ else List.map (fun p -> (p, 1.0 /. float_of_int k)) ps
+ in
+ Hashtbl.replace tbl lhs probs)
+ by_lhs;
+ tbl
+
+let canonical_string g =
+ let seen = Hashtbl.create 64 in
+ let order = ref [] in
+ let q = Queue.create () in
+ Queue.add g.start q;
+ Hashtbl.replace seen g.start ();
+ while not (Queue.is_empty q) do
+ let a = Queue.pop q in
+ order := a :: !order;
+ List.iter
+ (fun p ->
+ if String.equal p.lhs a then
+ List.iter
+ (function
+ | Nonterm x ->
+ if not (Hashtbl.mem seen x) then begin
+ Hashtbl.replace seen x ();
+ Queue.add x q
+ end
+ | Term _ -> ())
+ p.rhs)
+ g.productions
+ done;
+ let order = List.rev !order in
+ let idx = Hashtbl.create 64 in
+ List.iteri (fun i a -> Hashtbl.replace idx a i) order;
+ let name a =
+ match Hashtbl.find_opt idx a with
+ | Some i -> Printf.sprintf "N%d" i
+ | None -> a
+ in
+ let sym = function Term s -> "t:" ^ s | Nonterm s -> name s in
+ let pstr p =
+ Printf.sprintf "%s -> %s" (name p.lhs)
+ (String.concat " " (List.map sym p.rhs))
+ in
+ g.productions |> List.map pstr |> List.sort String.compare
+ |> String.concat "\n"
+
+let equal_structure a b = String.equal (canonical_string a) (canonical_string b)
+
+let copy_with g productions = { g with productions }
+
+let prods_of_list prods lhs =
+ List.filter (fun p -> String.equal p.lhs lhs) prods
+
+let dedupe prods =
+ let tbl = Hashtbl.create 64 in
+ List.iter
+ (fun p ->
+ match Hashtbl.find_opt tbl (p.lhs, p.rhs) with
+ | Some c -> Hashtbl.replace tbl (p.lhs, p.rhs) (c +. p.count)
+ | None -> Hashtbl.replace tbl (p.lhs, p.rhs) p.count)
+ prods;
+ Hashtbl.fold (fun (lhs, rhs) count acc -> { lhs; rhs; count } :: acc) tbl []
+
+let remove_unit_cycles prods =
+ let unit_edges = Hashtbl.create 64 in
+ List.iter
+ (fun (p : production) ->
+ match p.rhs with
+ | [ Nonterm b ] ->
+ let l =
+ match Hashtbl.find_opt unit_edges p.lhs with
+ | Some l -> l
+ | None -> []
+ in
+ Hashtbl.replace unit_edges p.lhs (b :: l)
+ | _ -> ())
+ prods;
+ let reaches a target =
+ let seen = Hashtbl.create 16 in
+ let rec go x =
+ if String.equal x target then true
+ else if Hashtbl.mem seen x then false
+ else begin
+ Hashtbl.replace seen x ();
+ List.exists go
+ (Option.value ~default:[] (Hashtbl.find_opt unit_edges x))
+ end
+ in
+ go a
+ in
+ List.filter
+ (fun (p : production) ->
+ match p.rhs with
+ | [ Nonterm b ] ->
+ not (String.equal p.lhs b) && not (reaches b p.lhs)
+ | _ -> true)
+ prods
+
+let productive_set prods =
+ let tbl = Hashtbl.create 64 in
+ let changed = ref true in
+ while !changed do
+ changed := false;
+ List.iter
+ (fun p ->
+ if not (Hashtbl.mem tbl p.lhs) then begin
+ let ok =
+ List.for_all
+ (function Term _ -> true | Nonterm x -> Hashtbl.mem tbl x)
+ p.rhs
+ in
+ if ok then begin
+ Hashtbl.replace tbl p.lhs ();
+ changed := true
+ end
+ end)
+ prods
+ done;
+ tbl
+
+let remove_nonproductive prods =
+ let prod = productive_set prods in
+ List.filter
+ (fun p ->
+ Hashtbl.mem prod p.lhs
+ && List.for_all
+ (function Term _ -> true | Nonterm x -> Hashtbl.mem prod x)
+ p.rhs)
+ prods
+
+let remove_unreachable start prods =
+ let tbl = Hashtbl.create 64 in
+ Hashtbl.replace tbl start ();
+ let changed = ref true in
+ while !changed do
+ changed := false;
+ List.iter
+ (fun p ->
+ if Hashtbl.mem tbl p.lhs then
+ List.iter
+ (function
+ | Nonterm x ->
+ if not (Hashtbl.mem tbl x) then begin
+ Hashtbl.replace tbl x ();
+ changed := true
+ end
+ | Term _ -> ())
+ p.rhs)
+ prods
+ done;
+ List.filter (fun p -> Hashtbl.mem tbl p.lhs) prods
+
+let derives_form g start target =
+ let n = List.length target in
+ let arr = Array.of_list target in
+ let key = symbol_key in
+ let d = Array.make_matrix (n + 1) (n + 1) SSet.empty in
+ for i = 0 to n - 1 do
+ d.(i).(i + 1) <- SSet.singleton (key arr.(i))
+ done;
+ let rhs_matches rhs i j =
+ let rec go rhs p =
+ match rhs with
+ | [] -> p = j
+ | [ last ] -> p < j && SSet.mem (key last) d.(p).(j)
+ | s :: rest ->
+ let rem = List.length rest in
+ let rec try_q q =
+ if q > j - rem then false
+ else if SSet.mem (key s) d.(p).(q) && go rest q then true
+ else try_q (q + 1)
+ in
+ try_q (p + 1)
+ in
+ List.length rhs > 0 && go rhs i
+ in
+ for len = 1 to n do
+ for i = 0 to n - len do
+ let j = i + len in
+ let changed = ref true in
+ while !changed do
+ changed := false;
+ List.iter
+ (fun (p : production) ->
+ let lhs_key = key (Nonterm p.lhs) in
+ if not (SSet.mem lhs_key d.(i).(j)) then begin
+ let ok =
+ match p.rhs with
+ | [ y ] -> SSet.mem (key y) d.(i).(j)
+ | _ -> rhs_matches p.rhs i j
+ in
+ if ok then begin
+ d.(i).(j) <- SSet.add lhs_key d.(i).(j);
+ changed := true
+ end
+ end)
+ g.productions
+ done
+ done
+ done;
+ SSet.mem (key (Nonterm start)) d.(0).(n)
+
+let prune_redundant g =
+ let prods = ref (dedupe g.productions) in
+ let changed = ref true in
+ while !changed do
+ changed := false;
+ let victim = ref None in
+ (try
+ List.iter
+ (fun (p : production) ->
+ let g' =
+ { g with productions = List.filter (fun q -> not (q == p)) !prods }
+ in
+ if derives_form g' p.lhs p.rhs then begin
+ victim := Some p;
+ raise Exit
+ end)
+ !prods
+ with Exit -> ());
+ match !victim with
+ | Some p ->
+ prods := List.filter (fun q -> not (q == p)) !prods;
+ changed := true
+ | None -> ()
+ done;
+ { g with productions = !prods }
+
+let normalize g =
+ let step prods =
+ prods |> dedupe |> remove_unit_cycles |> remove_nonproductive
+ |> remove_unreachable g.start
+ in
+ let p1 = step g.productions in
+ let p2 = step p1 in
+ let g1 = { g with productions = p2 } in
+ let g2 = prune_redundant g1 in
+ let g3 = step g2.productions in
+ { g with productions = g3 }
+
+let language_up_to g max_len =
+ let tbl = Hashtbl.create 256 in
+ let get x len =
+ Option.value ~default:[] (Hashtbl.find_opt tbl (x, len))
+ in
+ for len = 1 to max_len do
+
+ let changed = ref true in
+ while !changed do
+ changed := false;
+ List.iter
+ (fun x ->
+ let acc = ref (get x len) in
+ List.iter
+ (fun (p : production) ->
+ let rhs = p.rhs in
+ let m = List.length rhs in
+ let rec gen rhs l =
+ match rhs with
+ | [] -> if l = 0 then [ [] ] else []
+ | s :: rest ->
+ let out = ref [] in
+ for k = 1 to l - (List.length rest) do
+ let ss =
+ match s with
+ | Term a -> if k = 1 then [ [ a ] ] else []
+ | Nonterm y -> get y k
+ in
+ let rs = gen rest (l - k) in
+ List.iter
+ (fun a ->
+ List.iter (fun b -> out := (a @ b) :: !out) rs)
+ ss
+ done;
+ !out
+ in
+ if len >= m then acc := gen rhs len @ !acc)
+ (productions_of g x);
+ let acc = List.sort_uniq compare !acc in
+ if List.length acc > List.length (get x len) then begin
+ Hashtbl.replace tbl (x, len) acc;
+ changed := true
+ end)
+ (nonterminals g)
+ done
+ done;
+ let all = List.concat (List.init max_len (fun i -> get g.start (i + 1))) in
+ List.sort_uniq compare all
+
+let initial_grammar ?(start = "S") samples =
+ let term_nts = Hashtbl.create 64 in
+ let prods = ref [] in
+ List.iter
+ (fun (tokens, count) ->
+ let rhs =
+ List.map
+ (fun tok ->
+ match Hashtbl.find_opt term_nts tok with
+ | Some nt -> Nonterm nt
+ | None ->
+ let nt = "T_" ^ tok in
+ Hashtbl.replace term_nts tok nt;
+ prods := { lhs = nt; rhs = [ Term tok ]; count = 1.0 } :: !prods;
+ Nonterm nt)
+ tokens
+ in
+ prods := { lhs = start; rhs; count } :: !prods)
+ samples;
+ make ~start (List.rev !prods)
+
+let is_recursive g =
+ let deps = Hashtbl.create 64 in
+ List.iter
+ (fun x ->
+ let l =
+ List.concat_map
+ (fun p ->
+ List.filter_map
+ (function Nonterm y -> Some y | Term _ -> None)
+ p.rhs)
+ (productions_of g x)
+ in
+ Hashtbl.replace deps x l)
+ (nonterminals g);
+ List.exists
+ (fun a ->
+ let seen = Hashtbl.create 64 in
+ let rec go x =
+ List.exists
+ (fun y ->
+ if String.equal y a then true
+ else if Hashtbl.mem seen y then false
+ else begin
+ Hashtbl.replace seen y ();
+ go y
+ end)
+ (Option.value ~default:[] (Hashtbl.find_opt deps x))
+ in
+ go a)
+ (nonterminals g)
+
+let to_string ?(decimals = 3) g =
+ let probs = probabilities g in
+ let fmt p =
+ let pr =
+ match Hashtbl.find_opt probs p.lhs with
+ | None -> 0.0
+ | Some l -> (
+ match List.find_opt (fun (q, _) -> q == p) l with
+ | Some (_, x) -> x
+ | None -> 0.0)
+ in
+ Printf.sprintf "%-8s -> %-24s [%.*f]" p.lhs (rhs_to_string p.rhs) decimals
+ pr
+ in
+ let header =
+ Printf.sprintf "# start = %s, %d nonterminals, %d terminals, %d productions"
+ g.start
+ (List.length (nonterminals g))
+ (List.length (terminals g))
+ (List.length g.productions)
+ in
+ String.concat "\n" (header :: List.map fmt g.productions)
+
+let is_nonterminal_name s =
+ String.length s > 0
+ &&
+ let c = s.[0] in
+ (c >= 'A' && c <= 'Z') || c = '_' || c = '<'
+
+let find_sub s sub =
+ let n = String.length s and m = String.length sub in
+ let rec go i =
+ if i + m > n then None
+ else if String.equal (String.sub s i m) sub then Some i
+ else go (i + 1)
+ in
+ go 0
+
+let parse_start line =
+ match find_sub line "start =" with
+ | None -> None
+ | Some i ->
+ let rest =
+ String.sub line (i + 7) (String.length line - i - 7) |> String.trim
+ in
+ let stop =
+ match (String.index_opt rest ' ', String.index_opt rest ',') with
+ | Some a, Some b -> min a b
+ | Some a, None | None, Some a -> a
+ | None, None -> String.length rest
+ in
+ let name = String.sub rest 0 stop in
+ if String.length name = 0 then None else Some name
+
+let of_string text =
+ let lines = String.split_on_char '\n' text in
+ let prods = ref [] in
+ let start = ref None in
+ List.iter
+ (fun raw ->
+ let trimmed = String.trim raw in
+ if String.length trimmed > 0 && trimmed.[0] = '#' then
+ match parse_start trimmed with
+ | Some s -> start := Some s
+ | None -> ()
+ else begin
+ let line =
+ match String.index_opt raw '#' with
+ | Some i -> String.sub raw 0 i
+ | None -> raw
+ in
+ let toks =
+ String.split_on_char ' ' (String.trim line)
+ |> List.filter (fun s -> String.length s > 0)
+ in
+ match toks with
+ | [] -> ()
+ | lhs :: arrow :: rhs_toks when String.equal arrow "->" ->
+ let rhs_toks, prob =
+ match List.rev rhs_toks with
+ | last :: rest_rev
+ when String.length last > 2 && last.[0] = '[' -> (
+ match
+ float_of_string_opt
+ (String.sub last 1 (String.length last - 2))
+ with
+ | Some p -> (List.rev rest_rev, Some p)
+ | None -> (rhs_toks, None))
+ | _ -> (rhs_toks, None)
+ in
+ let rhs =
+ List.map
+ (fun s -> if is_nonterminal_name s then Nonterm s else Term s)
+ rhs_toks
+ in
+ if !start = None then start := Some lhs;
+ prods :=
+ { lhs; rhs; count = Option.value ~default:1.0 prob } :: !prods
+ | _ -> ()
+ end)
+ lines;
+ let start = Option.value ~default:"S" !start in
+ make ~start (List.rev !prods)
diff --git a/lib/parse.ml b/lib/parse.ml
new file mode 100644
index 0000000..3cfdb82
--- /dev/null
+++ b/lib/parse.ml
@@ -0,0 +1,355 @@
+type tree = Node of string * tree list | Leaf of string
+
+let rec tree_yield = function
+ | Leaf s -> [ s ]
+ | Node (_, children) -> List.concat_map tree_yield children
+
+type prep = {
+ prods : Grammar.production array;
+ mutable probs : float array;
+ by_lhs : (string, int list) Hashtbl.t;
+ lhs_list : string list;
+ start : string;
+}
+
+let prepare (g : Grammar.t) =
+ let prods : Grammar.production array = Array.of_list g.productions in
+ let n = Array.length prods in
+ let probs = Array.make n 0.0 in
+ let by_lhs = Hashtbl.create 64 in
+ Array.iteri
+ (fun i (p : Grammar.production) ->
+ let l =
+ match Hashtbl.find_opt by_lhs p.lhs with Some l -> l | None -> []
+ in
+ Hashtbl.replace by_lhs p.lhs (i :: l))
+ prods;
+ Hashtbl.iter
+ (fun _ idxs ->
+ let tot =
+ List.fold_left (fun s i -> s +. prods.(i).count) 0.0 idxs
+ in
+ let k = List.length idxs in
+ List.iter
+ (fun i ->
+ probs.(i) <-
+ (if tot > 0.0 then prods.(i).count /. tot
+ else 1.0 /. float_of_int k))
+ idxs)
+ by_lhs;
+ {
+ prods;
+ probs;
+ by_lhs;
+ lhs_list = Hashtbl.fold (fun k _ acc -> k :: acc) by_lhs [];
+ start = g.start;
+ }
+
+let prepare_uniform (g : Grammar.t) =
+ let prep = prepare g in
+ Array.fill prep.probs 0 (Array.length prep.probs) 0.0;
+ Hashtbl.iter
+ (fun _ idxs ->
+ let k = List.length idxs in
+ if k > 0 then
+ List.iter (fun i -> prep.probs.(i) <- 1.0 /. float_of_int k) idxs)
+ prep.by_lhs;
+ prep
+
+let prob_of_prod prep (p : Grammar.production) =
+ let idxs =
+ Option.value ~default:[] (Hashtbl.find_opt prep.by_lhs p.lhs)
+ in
+ let rec find = function
+ | [] -> 0.0
+ | i :: rest ->
+ if Grammar.equal_rhs prep.prods.(i).rhs p.rhs then prep.probs.(i)
+ else find rest
+ in
+ find idxs
+
+let inside prep tokens =
+ let n = List.length tokens in
+ let w = Array.of_list tokens in
+ let imemo = Hashtbl.create 1024 in
+ let smemo = Hashtbl.create 1024 in
+ let rec inside_nt nt i j =
+ if j <= i then 0.0
+ else
+ match Hashtbl.find_opt imemo (nt, i, j) with
+ | Some v -> v
+ | None ->
+ let idxs =
+ Option.value ~default:[] (Hashtbl.find_opt prep.by_lhs nt)
+ in
+ let v =
+ List.fold_left
+ (fun acc pi ->
+ acc
+ +. (prep.probs.(pi) *. seq_inside prep.prods.(pi).rhs i j))
+ 0.0 idxs
+ in
+ Hashtbl.replace imemo (nt, i, j) v;
+ v
+ and seq_inside rhs i j =
+ match Hashtbl.find_opt smemo (rhs, i, j) with
+ | Some v -> v
+ | None ->
+ let v =
+ match rhs with
+ | [] -> if i = j then 1.0 else 0.0
+ | Grammar.Term a :: rest ->
+ if i < n && String.equal w.(i) a then
+ seq_inside rest (i + 1) j
+ else 0.0
+ | Grammar.Nonterm x :: rest ->
+
+ let acc = ref 0.0 in
+ for k = i + 1 to j - List.length rest do
+ let a = inside_nt x i k in
+ if a > 0.0 then acc := !acc +. (a *. seq_inside rest k j)
+ done;
+ !acc
+ in
+ Hashtbl.replace smemo (rhs, i, j) v;
+ v
+ in
+ inside_nt prep.start 0 n
+
+let log_inside prep tokens =
+ let p = inside prep tokens in
+ if p > 0.0 then Some (log p) else None
+
+let viterbi prep tokens =
+ let n = List.length tokens in
+ let w = Array.of_list tokens in
+ let vmemo = Hashtbl.create 1024 in
+ let smemo = Hashtbl.create 1024 in
+ let rec vnt nt i j =
+ if j <= i then (0.0, None)
+ else
+ match Hashtbl.find_opt vmemo (nt, i, j) with
+ | Some v -> v
+ | None ->
+ let idxs =
+ Option.value ~default:[] (Hashtbl.find_opt prep.by_lhs nt)
+ in
+ let bestp = ref 0.0 and bestt = ref None in
+ List.iter
+ (fun pi ->
+ let p = prep.prods.(pi) in
+ let sp, sch = seq_vit p.rhs i j in
+ if sp > 0.0 then begin
+ let v = prep.probs.(pi) *. sp in
+ if v > !bestp then begin
+ bestp := v;
+ bestt :=
+ (match sch with
+ | Some ch -> Some (Node (nt, ch))
+ | None -> None)
+ end
+ end)
+ idxs;
+ let res = (!bestp, !bestt) in
+ Hashtbl.replace vmemo (nt, i, j) res;
+ res
+ and seq_vit rhs i j =
+ match Hashtbl.find_opt smemo (rhs, i, j) with
+ | Some v -> v
+ | None ->
+ let v =
+ match rhs with
+ | [] -> if i = j then (1.0, Some []) else (0.0, None)
+ | Grammar.Term a :: rest ->
+ if i < n && String.equal w.(i) a then (
+ match seq_vit rest (i + 1) j with
+ | p, Some ts -> (p, Some (Leaf a :: ts))
+ | z -> z)
+ else (0.0, None)
+ | Grammar.Nonterm x :: rest ->
+ let bestp = ref 0.0 and bestch = ref None in
+ for k = i + 1 to j - List.length rest do
+ let px, tx = vnt x i k in
+ if px > 0.0 then begin
+ let pr, tr = seq_vit rest k j in
+ if pr > 0.0 then begin
+ let v = px *. pr in
+ if v > !bestp then begin
+ bestp := v;
+ bestch :=
+ (match (tx, tr) with
+ | Some tx, Some tr -> Some (tx :: tr)
+ | _ -> None)
+ end
+ end
+ end
+ done;
+ (!bestp, !bestch)
+ in
+ Hashtbl.replace smemo (rhs, i, j) v;
+ v
+ in
+ vnt prep.start 0 n
+
+let tree_to_string ?(indent = 0) tree =
+ let rec go indent tree =
+ let pad = String.make (2 * indent) ' ' in
+ match tree with
+ | Leaf s -> pad ^ s
+ | Node (nt, children) ->
+ pad ^ nt ^ "\n"
+ ^ String.concat "\n" (List.map (go (indent + 1)) children)
+ in
+ go indent tree
+
+let expected_counts prep tokens =
+ let n = List.length tokens in
+ let w = Array.of_list tokens in
+ let imemo = Hashtbl.create 1024 in
+ let smemo = Hashtbl.create 1024 in
+ let rec inside_nt nt i j =
+ if j <= i then 0.0
+ else
+ match Hashtbl.find_opt imemo (nt, i, j) with
+ | Some v -> v
+ | None ->
+ let idxs =
+ Option.value ~default:[] (Hashtbl.find_opt prep.by_lhs nt)
+ in
+ let v =
+ List.fold_left
+ (fun acc pi ->
+ acc
+ +. (prep.probs.(pi) *. seq_inside prep.prods.(pi).rhs i j))
+ 0.0 idxs
+ in
+ Hashtbl.replace imemo (nt, i, j) v;
+ v
+ and seq_inside rhs i j =
+ match Hashtbl.find_opt smemo (rhs, i, j) with
+ | Some v -> v
+ | None ->
+ let v =
+ match rhs with
+ | [] -> if i = j then 1.0 else 0.0
+ | Grammar.Term a :: rest ->
+ if i < n && String.equal w.(i) a then seq_inside rest (i + 1) j
+ else 0.0
+ | Grammar.Nonterm x :: rest ->
+
+ let acc = ref 0.0 in
+ for k = i + 1 to j - List.length rest do
+ let a = inside_nt x i k in
+ if a > 0.0 then acc := !acc +. (a *. seq_inside rest k j)
+ done;
+ !acc
+ in
+ Hashtbl.replace smemo (rhs, i, j) v;
+ v
+ in
+ let ptotal = inside_nt prep.start 0 n in
+ let np = Array.length prep.prods in
+ if ptotal <= 0.0 then (Array.make np 0.0, 0.0)
+ else begin
+ let omemo = Hashtbl.create 1024 in
+ let get_o nt i j =
+ Option.value ~default:0.0 (Hashtbl.find_opt omemo (nt, i, j))
+ in
+ let add_o nt i j v =
+ Hashtbl.replace omemo (nt, i, j) (get_o nt i j +. v)
+ in
+ add_o prep.start 0 n 1.0;
+ for len = n downto 1 do
+ for p = 0 to n - len do
+ let q = p + len in
+ List.iter
+ (fun b ->
+ let ob = get_o b p q in
+ if ob > 0.0 then
+ List.iter
+ (fun pi ->
+ let prod = prep.prods.(pi) in
+ if String.equal prod.lhs b then begin
+ let prob = prep.probs.(pi) in
+ let arr = Array.of_list prod.rhs in
+ let m = Array.length arr in
+ Array.iteri
+ (fun k s ->
+ match s with
+ | Grammar.Nonterm x ->
+ let left =
+ Array.to_list (Array.sub arr 0 k)
+ in
+ let right =
+ Array.to_list
+ (Array.sub arr (k + 1) (m - k - 1))
+ in
+ for i = p to q do
+ let lw = seq_inside left p i in
+ if lw > 0.0 then
+ for j = i to q do
+ let rw = seq_inside right j q in
+ if rw > 0.0 then
+ add_o x i j
+ (ob *. prob *. lw *. rw)
+ done
+ done
+ | Grammar.Term _ -> ())
+ arr
+ end)
+ (Option.value ~default:[]
+ (Hashtbl.find_opt prep.by_lhs b)))
+ prep.lhs_list
+ done
+ done;
+ let counts = Array.make np 0.0 in
+ Array.iteri
+ (fun pi (prod : Grammar.production) ->
+ let prob = prep.probs.(pi) in
+ let s = ref 0.0 in
+ for i = 0 to n do
+ for j = i to n do
+ let o = get_o prod.lhs i j in
+ if o > 0.0 then
+ s := !s +. (o *. seq_inside prod.rhs i j)
+ done
+ done;
+ counts.(pi) <- prob *. !s /. ptotal)
+ prep.prods;
+ (counts, ptotal)
+ end
+
+let fit_em prep (corpus : Corpus.t) ~iters ~tol =
+ let np = Array.length prep.prods in
+ let last_ll = ref neg_infinity in
+ let ll = ref 0.0 in
+ let counts = ref (Array.make np 0.0) in
+ (try
+ for _ = 1 to iters do
+ let acc = Array.make np 0.0 in
+ ll := 0.0;
+ List.iter
+ (fun s ->
+ let cnt, p = expected_counts prep s.Corpus.tokens in
+ if p > 0.0 then begin
+ Array.iteri
+ (fun i x -> acc.(i) <- acc.(i) +. (s.Corpus.count *. x))
+ cnt;
+ ll := !ll +. (s.Corpus.count *. log p)
+ end)
+ corpus;
+ Hashtbl.iter
+ (fun _ idxs ->
+ let tot = List.fold_left (fun s i -> s +. acc.(i)) 0.0 idxs in
+ if tot > 0.0 then
+ List.iter (fun i -> prep.probs.(i) <- acc.(i) /. tot) idxs)
+ prep.by_lhs;
+ counts := acc;
+ if
+ !last_ll <> neg_infinity
+ && abs_float (!ll -. !last_ll) < tol
+ then raise Exit;
+ last_ll := !ll
+ done
+ with Exit -> ());
+ (!counts, !ll)
diff --git a/lib/scoring.ml b/lib/scoring.ml
new file mode 100644
index 0000000..9712b57
--- /dev/null
+++ b/lib/scoring.ml
@@ -0,0 +1,278 @@
+type score_mode = Marginal | Maximum_likelihood | Variational
+
+type chunk_occurrence = All_occurrences | First_occurrence
+
+type config = {
+ alpha : float;
+ prior_weight : float;
+ em_iters : int;
+ em_tol : float;
+ max_chunk : int;
+ mode : score_mode;
+ chunk_occurrence : chunk_occurrence;
+}
+
+let default_config =
+ {
+ alpha = 0.1;
+ prior_weight = 1.0;
+ em_iters = 8;
+ em_tol = 1e-7;
+ max_chunk = 2;
+ mode = Marginal;
+ chunk_occurrence = All_occurrences;
+ }
+
+let rec log_gamma x =
+ let g = 7.0 in
+ let c =
+ [|
+ 0.99999999999980993;
+ 676.5203681218851;
+ -1259.1392167224028;
+ 771.32342877765313;
+ -176.61502916214059;
+ 12.507343278686905;
+ -0.13857109526572012;
+ 9.9843695780195716e-6;
+ 1.5056327351493116e-7;
+ |]
+ in
+ if x < 0.5 then
+ log (Float.pi /. sin (Float.pi *. x)) -. log_gamma (1.0 -. x)
+ else begin
+ let x = x -. 1.0 in
+ let a = ref c.(0) in
+ for i = 1 to 8 do
+ a := !a +. (c.(i) /. (x +. float_of_int i))
+ done;
+ let t = x +. g +. 0.5 in
+ (0.5 *. log (2.0 *. Float.pi))
+ +. ((x +. 0.5) *. log t)
+ -. t +. log !a
+ end
+
+let rec digamma x =
+ if x < 6.0 then digamma (x +. 1.0) -. (1.0 /. x)
+ else
+ let inv = 1.0 /. x in
+ let inv2 = inv *. inv in
+ log x -. (0.5 *. inv)
+ -. (inv2
+ *. (1.0 /. 12.0
+ -. (inv2
+ *. (1.0 /. 120.0
+ -. (inv2
+ *. (1.0 /. 252.0
+ -. (inv2 *. (1.0 /. 240.0 -. (inv2 *. (1.0 /. 132.0))))))))))
+
+let bits x = if x <= 0.0 then 0.0 else log x /. log 2.0
+
+let description_length_bits (g : Grammar.t) =
+ let nts = Grammar.nonterminals g in
+ let ts = Grammar.terminals g in
+ let nn = float_of_int (List.length nts) in
+ let nt_terms = float_of_int (List.length ts) in
+ let bnt = bits nn in
+ let bt = bits nt_terms in
+ let total = ref (bits nn +. bits nt_terms) in
+ List.iter
+ (fun a ->
+ let ps = Grammar.productions_of g a in
+ total := !total +. bits (float_of_int (List.length ps));
+ List.iter
+ (fun (p : Grammar.production) ->
+ total := !total +. bits (float_of_int (List.length p.rhs));
+ List.iter
+ (function
+ | Grammar.Nonterm _ -> total := !total +. bnt
+ | Grammar.Term _ -> total := !total +. bt)
+ p.rhs)
+ ps)
+ nts;
+ !total
+
+let structural_logprior ?(config = default_config) (g : Grammar.t) =
+ -.config.prior_weight *. description_length_bits g *. log 2.0
+
+let marginal_loglik ?(config = default_config) (g : Grammar.t) (corpus : Corpus.t) =
+
+ let uprep = Parse.prepare_uniform g in
+ if not (List.for_all (fun s -> Parse.inside uprep s.Corpus.tokens > 0.0) corpus)
+ then (neg_infinity, neg_infinity, [||], Parse.prepare g)
+ else begin
+ let prep = Parse.prepare g in
+ ignore (Parse.fit_em prep corpus ~iters:config.em_iters ~tol:config.em_tol);
+
+ let ll =
+ List.fold_left
+ (fun acc s ->
+ let p = Parse.inside prep s.Corpus.tokens in
+ if p > 0.0 then acc +. (s.Corpus.count *. log p) else neg_infinity)
+ 0.0 corpus
+ in
+ if ll = neg_infinity then (neg_infinity, neg_infinity, [||], prep)
+ else begin
+
+ let np = Array.length prep.prods in
+ let counts = Array.make np 0.0 in
+ List.iter
+ (fun s ->
+ let cnt, _ = Parse.expected_counts prep s.Corpus.tokens in
+ Array.iteri
+ (fun i x -> counts.(i) <- counts.(i) +. (s.Corpus.count *. x))
+ cnt)
+ corpus;
+ let dirichlet = ref 0.0 and ml_count = ref 0.0 in
+ Hashtbl.iter
+ (fun _ idxs ->
+ let k = List.length idxs in
+ if k > 0 then begin
+ let n = List.fold_left (fun s i -> s +. counts.(i)) 0.0 idxs in
+ if n > 0.0 then begin
+ dirichlet :=
+ !dirichlet
+ +. (log_gamma (float_of_int k *. config.alpha)
+ -. log_gamma (n +. (float_of_int k *. config.alpha)));
+ List.iter
+ (fun i ->
+ if counts.(i) > 0.0 then begin
+ dirichlet :=
+ !dirichlet
+ +. (log_gamma (counts.(i) +. config.alpha)
+ -. log_gamma config.alpha);
+ ml_count :=
+ !ml_count +. (counts.(i) *. log (counts.(i) /. n))
+ end)
+ idxs
+ end
+ end)
+ prep.by_lhs;
+ let occam = !dirichlet -. !ml_count in
+ (ll +. occam, ll, counts, prep)
+ end
+ end
+
+let variational_loglik ?(config = default_config) (g : Grammar.t) (corpus : Corpus.t) =
+ let prep = Parse.prepare g in
+ let np = Array.length prep.prods in
+ let beta = Array.make np config.alpha in
+ let kl = ref 0.0 in
+ let set_theta_bar () =
+ Hashtbl.iter
+ (fun _ idxs ->
+ let bsum = List.fold_left (fun s i -> s +. beta.(i)) 0.0 idxs in
+ let bpsi = digamma bsum in
+ List.iter (fun i -> prep.probs.(i) <- exp (digamma beta.(i) -. bpsi)) idxs)
+ prep.by_lhs
+ in
+ let update_kl () =
+ kl := 0.0;
+ Hashtbl.iter
+ (fun _ idxs ->
+ let k = List.length idxs in
+ if k > 0 then begin
+ let bsum = List.fold_left (fun s i -> s +. beta.(i)) 0.0 idxs in
+ let bpsi = digamma bsum in
+ let log_b_alpha =
+ (float_of_int k *. log_gamma config.alpha)
+ -. log_gamma (float_of_int k *. config.alpha)
+ in
+ let log_b_beta =
+ List.fold_left (fun s i -> s +. log_gamma beta.(i)) 0.0 idxs
+ -. log_gamma bsum
+ in
+ let term =
+ List.fold_left
+ (fun s i ->
+ s +. ((beta.(i) -. config.alpha) *. (digamma beta.(i) -. bpsi)))
+ 0.0 idxs
+ in
+ kl := !kl +. (log_b_alpha -. log_b_beta +. term)
+ end)
+ prep.by_lhs
+ in
+ for _ = 1 to max 1 config.em_iters do
+ set_theta_bar ();
+ let acc = Array.make np 0.0 in
+ List.iter
+ (fun s ->
+ let cnt, _ = Parse.expected_counts prep s.Corpus.tokens in
+ Array.iteri
+ (fun i x -> acc.(i) <- acc.(i) +. (s.Corpus.count *. x))
+ cnt)
+ corpus;
+ for i = 0 to np - 1 do
+ beta.(i) <- config.alpha +. acc.(i)
+ done;
+ update_kl ()
+ done;
+ set_theta_bar ();
+ let log_z = ref 0.0 in
+ List.iter
+ (fun s ->
+ let z = Parse.inside prep s.Corpus.tokens in
+ if z > 0.0 then log_z := !log_z +. (s.Corpus.count *. log z)
+ else log_z := neg_infinity)
+ corpus;
+ !log_z -. !kl
+
+let posterior ?(config = default_config) g corpus =
+ let prior = structural_logprior ~config g in
+ match config.mode with
+ | Variational -> prior +. variational_loglik ~config g corpus
+ | _ ->
+ let marg, ll, _, _ = marginal_loglik ~config g corpus in
+ (match config.mode with
+ | Marginal -> prior +. marg
+ | Maximum_likelihood -> prior +. ll
+ | Variational -> assert false)
+
+type details = {
+ prior : float;
+ marginal : float;
+ ml_loglik : float;
+ posterior : float;
+ dl_bits : float;
+ num_nonterminals : int;
+ num_productions : int;
+}
+
+let details ?(config = default_config) (g : Grammar.t) (corpus : Corpus.t) =
+ let marg, ll, _, _ = marginal_loglik ~config g corpus in
+ let prior = structural_logprior ~config g in
+ let dl = description_length_bits g in
+ let likelihood =
+ match config.mode with
+ | Marginal -> marg
+ | Maximum_likelihood -> ll
+ | Variational -> variational_loglik ~config g corpus
+ in
+ {
+ prior;
+ marginal = likelihood;
+ ml_loglik = ll;
+ posterior = prior +. likelihood;
+ dl_bits = dl;
+ num_nonterminals = List.length (Grammar.nonterminals g);
+ num_productions = List.length g.Grammar.productions;
+ }
+
+let heldout ?(config = default_config) (g : Grammar.t) ~train (test : Corpus.t) =
+ let prep = Parse.prepare g in
+ ignore (Parse.fit_em prep train ~iters:config.em_iters ~tol:config.em_tol);
+ let total = ref 0.0 in
+ let ntok = ref 0.0 in
+ let unparseable = ref 0 in
+ List.iter
+ (fun s ->
+ let p = Parse.inside prep s.Corpus.tokens in
+ if p > 0.0 then begin
+ total := !total +. (s.Corpus.count *. log p);
+ ntok := !ntok +. (s.Corpus.count *. float_of_int (List.length s.Corpus.tokens))
+ end
+ else incr unparseable)
+ test;
+ if !ntok > 0.0 then
+ (!total /. !ntok, !total, !ntok, !unparseable)
+ else (neg_infinity, !total, 0.0, !unparseable)
diff --git a/lib/search.ml b/lib/search.ml
new file mode 100644
index 0000000..614b574
--- /dev/null
+++ b/lib/search.ml
@@ -0,0 +1,250 @@
+type step = {
+ transform : Transform.t;
+ parent : Grammar.t;
+ result : Grammar.t;
+ parent_score : float;
+ score : float;
+}
+
+type proposal = {
+ iteration : int;
+ transform : Transform.t;
+ parent_key : string;
+ score : float;
+ accepted : bool;
+}
+
+type entry = { g : Grammar.t; score : float; hist : step list }
+
+type result = {
+ best : Grammar.t;
+ best_score : float;
+ initial_score : float;
+ steps : step list;
+ proposals : proposal list;
+ iterations : int;
+}
+
+let key_of = Grammar.canonical_string
+
+let transform_key (s : step) =
+ (key_of s.parent, Transform.describe s.transform)
+
+let run ~config ~beam_width ~max_steps ~patience ?(on_step = fun _ _ -> ()) g0
+ corpus =
+ let score g = Scoring.posterior ~config g corpus in
+ let s0 = score g0 in
+ let beam = ref [ { g = g0; score = s0; hist = [] } ] in
+ let best = ref { g = g0; score = s0; hist = [] } in
+ let proposals = ref [] in
+ let accepted = Hashtbl.create 64 in
+ let no_improve = ref 0 in
+ let iter = ref 0 in
+ let stop = ref false in
+ while (not !stop) && !iter < max_steps do
+ incr iter;
+ let cands =
+ List.concat_map
+ (fun e ->
+ List.filter_map
+ (fun t ->
+ let g' = Grammar.normalize (Transform.apply ~occurrence:config.chunk_occurrence e.g t) in
+ let s' = score g' in
+ if s' = neg_infinity then None
+ else
+ let step =
+ {
+ transform = t;
+ parent = e.g;
+ result = g';
+ parent_score = e.score;
+ score = s';
+ }
+ in
+ Some { g = g'; score = s'; hist = step :: e.hist })
+ (Transform.candidates e.g config))
+ !beam
+ in
+ let tbl = Hashtbl.create 128 in
+ List.iter
+ (fun e ->
+ let k = key_of e.g in
+ match Hashtbl.find_opt tbl k with
+ | Some e0 when e0.score >= e.score -> ()
+ | _ -> Hashtbl.replace tbl k e)
+ cands;
+ let cands = Hashtbl.fold (fun _ e acc -> e :: acc) tbl [] in
+ List.iter
+ (fun e ->
+ match e.hist with
+ | step :: _ ->
+ proposals :=
+ {
+ iteration = !iter;
+ transform = step.transform;
+ parent_key = key_of step.parent;
+ score = e.score;
+ accepted = false;
+ }
+ :: !proposals
+ | [] -> ())
+ cands;
+ let sorted = List.sort (fun a b -> compare b.score a.score) cands in
+ let newbeam = List.filteri (fun i _ -> i < beam_width) sorted in
+ (match newbeam with
+ | e :: _ when e.score > !best.score +. 1e-9 ->
+ best := e;
+ (match e.hist with
+ | step :: _ ->
+ Hashtbl.replace accepted (transform_key step) ();
+ on_step !iter step
+ | [] -> ());
+ no_improve := 0
+ | _ -> incr no_improve);
+ beam := newbeam;
+ if newbeam = [] then stop := true;
+ if !no_improve >= patience then stop := true
+ done;
+ let steps = List.rev !best.hist in
+ let proposals =
+ List.rev !proposals
+ |> List.map (fun p ->
+ {
+ p with
+ accepted =
+ Hashtbl.mem accepted
+ (p.parent_key, Transform.describe p.transform);
+ })
+ in
+ {
+ best = !best.g;
+ best_score = !best.score;
+ initial_score = s0;
+ steps;
+ proposals;
+ iterations = !iter;
+ }
+
+let run_best_first ~config ~frontier_cap ~max_expansions ~patience g0 corpus =
+ let score g = Scoring.posterior ~config g corpus in
+ let s0 = score g0 in
+ let frontier = ref [ { g = g0; score = s0; hist = [] } ] in
+ let best = ref { g = g0; score = s0; hist = [] } in
+ let proposals = ref [] in
+ let accepted = Hashtbl.create 64 in
+ let no_improve = ref 0 in
+ let expansions = ref 0 in
+ let iter = ref 0 in
+ let dedupe entries =
+ let tbl = Hashtbl.create 256 in
+ List.iter
+ (fun e ->
+ let k = key_of e.g in
+ match Hashtbl.find_opt tbl k with
+ | Some e0 when e0.score >= e.score -> ()
+ | _ -> Hashtbl.replace tbl k e)
+ entries;
+ Hashtbl.fold (fun _ e acc -> e :: acc) tbl []
+ in
+ while
+ !frontier <> [] && !expansions < max_expansions && !no_improve < patience
+ do
+ incr iter;
+ let sorted = List.sort (fun a b -> compare b.score a.score) !frontier in
+ let e = List.hd sorted in
+ frontier := List.tl sorted;
+ incr expansions;
+ let cands =
+ List.filter_map
+ (fun t ->
+ let g' = Grammar.normalize (Transform.apply ~occurrence:config.chunk_occurrence e.g t) in
+ let s' = score g' in
+ if s' = neg_infinity then None
+ else
+ Some
+ {
+ g = g';
+ score = s';
+ hist =
+ {
+ transform = t;
+ parent = e.g;
+ result = g';
+ parent_score = e.score;
+ score = s';
+ }
+ :: e.hist;
+ })
+ (Transform.candidates e.g config)
+ in
+ List.iter
+ (fun c ->
+ match c.hist with
+ | step :: _ ->
+ proposals :=
+ {
+ iteration = !iter;
+ transform = step.transform;
+ parent_key = key_of step.parent;
+ score = c.score;
+ accepted = false;
+ }
+ :: !proposals
+ | [] -> ())
+ cands;
+ if e.score > !best.score +. 1e-9 then begin
+ best := e;
+ (match e.hist with
+ | step :: _ -> Hashtbl.replace accepted (transform_key step) ()
+ | [] -> ());
+ no_improve := 0
+ end
+ else incr no_improve;
+ frontier := cands @ !frontier |> dedupe;
+ let sorted = List.sort (fun a b -> compare b.score a.score) !frontier in
+ frontier := List.filteri (fun i _ -> i < frontier_cap) sorted
+ done;
+ let steps = List.rev !best.hist in
+ let proposals =
+ List.rev !proposals
+ |> List.map (fun p ->
+ {
+ p with
+ accepted =
+ Hashtbl.mem accepted
+ (p.parent_key, Transform.describe p.transform);
+ })
+ in
+ {
+ best = !best.g;
+ best_score = !best.score;
+ initial_score = s0;
+ steps;
+ proposals;
+ iterations = !iter;
+ }
+
+let proposals_to_tsv result =
+ let header = "iteration\tparent_key\ttransform\tscore\taccepted" in
+ let rows =
+ List.map
+ (fun p ->
+ Printf.sprintf "%d\t%s\t%s\t%.10g\t%b" p.iteration
+ (String.concat "\\n" (String.split_on_char '\n' p.parent_key))
+ (Transform.describe p.transform)
+ p.score p.accepted)
+ result.proposals
+ in
+ String.concat "\n" (header :: rows)
+
+let steps_to_tsv result =
+ let header = "step\tparent_score\tscore\tdelta\ttransform" in
+ let rows =
+ List.mapi
+ (fun i (s : step) ->
+ Printf.sprintf "%d\t%.10g\t%.10g\t%.10g\t%s" (i + 1) s.parent_score
+ s.score (s.score -. s.parent_score)
+ (Transform.describe s.transform))
+ result.steps
+ in
+ String.concat "\n" (header :: rows)
diff --git a/lib/transform.ml b/lib/transform.ml
new file mode 100644
index 0000000..2aa4a7c
--- /dev/null
+++ b/lib/transform.ml
@@ -0,0 +1,193 @@
+type t = Merge of string * string | Chunk of Grammar.symbol list
+
+let describe = function
+ | Merge (a, b) -> Printf.sprintf "merge %s,%s" a b
+ | Chunk seq ->
+ Printf.sprintf "chunk (%s)"
+ (String.concat " " (List.map Grammar.symbol_to_string seq))
+
+let starts_with seq lst =
+ let m = List.length seq in
+ let rec go i seq lst =
+ if i = 0 then true
+ else
+ match (seq, lst) with
+ | [], _ -> true
+ | _, [] -> false
+ | x :: xs, y :: ys -> Grammar.equal_symbol x y && go (i - 1) xs ys
+ in
+ List.length lst >= m && go m seq lst
+
+let drop n lst =
+ let rec go n lst = if n <= 0 then lst else match lst with [] -> [] | _ :: r -> go (n - 1) r in
+ go n lst
+
+let replace_all seq name rhs =
+ let m = List.length seq in
+ let rec go acc lst count =
+ match lst with
+ | _ when starts_with seq lst ->
+ go (Grammar.Nonterm name :: acc) (drop m lst) (count + 1)
+ | [] -> (List.rev acc, count)
+ | x :: rest -> go (x :: acc) rest count
+ in
+ go [] rhs 0
+
+let replace_first seq name rhs =
+ let m = List.length seq in
+ let rec go acc lst =
+ match lst with
+ | _ when starts_with seq lst -> (List.rev_append acc (Grammar.Nonterm name :: drop m lst), 1)
+ | [] -> (List.rev acc, 0)
+ | x :: rest -> go (x :: acc) rest
+ in
+ go [] rhs
+
+let fresh_name (g : Grammar.t) _base =
+ let existing = Grammar.nonterminals g in
+ let candidates = [ "X"; "Y"; "Z"; "W"; "V"; "U"; "P"; "Q"; "R" ] in
+ match List.find_opt (fun c -> not (List.mem c existing)) candidates with
+ | Some c -> c
+ | None ->
+ let rec loop i =
+ let c = "X" ^ string_of_int i in
+ if List.mem c existing then loop (i + 1) else c
+ in
+ loop 1
+
+let apply_merge (g : Grammar.t) a b =
+ let survivor, removed =
+ if String.equal a g.Grammar.start then (a, b)
+ else if String.equal b g.Grammar.start then (b, a)
+ else if String.compare a b <= 0 then (a, b)
+ else (b, a)
+ in
+ let map_sym = function
+ | Grammar.Nonterm x when String.equal x removed -> Grammar.Nonterm survivor
+ | s -> s
+ in
+ let prods =
+ List.map
+ (fun (p : Grammar.production) ->
+ let lhs = if String.equal p.lhs removed then survivor else p.lhs in
+ { p with lhs; rhs = List.map map_sym p.rhs })
+ g.Grammar.productions
+ in
+ Grammar.copy_with g prods
+
+let apply_chunk ?(occurrence = Scoring.All_occurrences) (g : Grammar.t) seq =
+ if List.length seq < 2 then g
+ else begin
+ let name = fresh_name g "X" in
+ let replace =
+ match occurrence with
+ | Scoring.All_occurrences -> replace_all
+ | Scoring.First_occurrence -> replace_first
+ in
+ let total = ref 0 in
+ let prods =
+ List.map
+ (fun (p : Grammar.production) ->
+ let rhs, n = replace seq name p.rhs in
+ total := !total + n;
+ { p with rhs })
+ g.Grammar.productions
+ in
+ let newp =
+ {
+ Grammar.lhs = name;
+ rhs = seq;
+ count = float_of_int (max 1 !total);
+ }
+ in
+ Grammar.copy_with g (prods @ [ newp ])
+ end
+
+let apply ?occurrence (g : Grammar.t) = function
+ | Merge (a, b) -> apply_merge g a b
+ | Chunk seq -> apply_chunk ?occurrence g seq
+
+let candidates (g : Grammar.t) (config : Scoring.config) =
+ let nts = Grammar.nonterminals g in
+ let merges =
+ let rec pairs = function
+ | [] -> []
+ | x :: rest -> List.map (fun y -> Merge (x, y)) rest @ pairs rest
+ in
+ pairs nts
+ in
+ let counts = Hashtbl.create 256 in
+ let order = ref [] in
+ List.iter
+ (fun (p : Grammar.production) ->
+ let arr = Array.of_list p.rhs in
+ let n = Array.length arr in
+ for len = 2 to min config.Scoring.max_chunk n do
+ for i = 0 to n - len do
+ let sub = Array.to_list (Array.sub arr i len) in
+ match Hashtbl.find_opt counts sub with
+ | Some c -> Hashtbl.replace counts sub (c + 1)
+ | None ->
+ Hashtbl.replace counts sub 1;
+ order := sub :: !order
+ done
+ done)
+ g.Grammar.productions;
+
+ let chunks =
+ List.rev !order
+ |> List.filter (fun k -> Hashtbl.find counts k >= 1)
+ |> List.map (fun seq -> Chunk seq)
+ in
+ merges @ chunks
+
+type change = {
+ before_nts : string list;
+ after_nts : string list;
+ before_prods : Grammar.production list;
+ after_prods : Grammar.production list;
+}
+
+let change_of (before : Grammar.t) t (after : Grammar.t) =
+ match t with
+ | Merge (a, b) ->
+ let survivor =
+ if String.equal a before.Grammar.start then a
+ else if String.equal b before.Grammar.start then b
+ else if String.compare a b <= 0 then a
+ else b
+ in
+ {
+ before_nts = [ a; b ];
+ after_nts = [ survivor ];
+ before_prods =
+ List.filter
+ (fun p -> String.equal p.Grammar.lhs a || String.equal p.Grammar.lhs b)
+ before.Grammar.productions;
+ after_prods =
+ List.filter
+ (fun p -> String.equal p.Grammar.lhs survivor)
+ after.Grammar.productions;
+ }
+ | Chunk seq ->
+ let name = fresh_name before "X" in
+ let affected_before =
+ List.filter (fun p -> starts_with seq p.Grammar.rhs) before.Grammar.productions
+ in
+ let affected_after =
+ List.filter
+ (fun p ->
+ String.equal p.Grammar.lhs name
+ || List.exists
+ (function
+ | Grammar.Nonterm x -> String.equal x name
+ | Grammar.Term _ -> false)
+ p.Grammar.rhs)
+ after.Grammar.productions
+ in
+ {
+ before_nts = name :: List.map (fun p -> p.Grammar.lhs) affected_before;
+ after_nts = [ name ];
+ before_prods = affected_before;
+ after_prods = affected_after;
+ }
diff --git a/meta/paper-and-induced.svg b/meta/paper-and-induced.svg
new file mode 100644
index 0000000..b56c2e6
--- /dev/null
+++ b/meta/paper-and-induced.svg
@@ -0,0 +1,9380 @@
+
\ No newline at end of file
diff --git a/test/dune b/test/dune
new file mode 100644
index 0000000..1351dd1
--- /dev/null
+++ b/test/dune
@@ -0,0 +1,4 @@
+(test
+ (name test_scfg)
+ (flags (:standard -warn-error -a))
+ (libraries scfg))
diff --git a/test/test_scfg.ml b/test/test_scfg.ml
new file mode 100644
index 0000000..a3cc896
--- /dev/null
+++ b/test/test_scfg.ml
@@ -0,0 +1,317 @@
+open Scfg
+
+let failures = ref 0
+
+let check name cond =
+ if cond then Printf.printf "ok %s\n" name
+ else begin
+ incr failures;
+ Printf.printf "FAIL %s\n" name
+ end;
+ flush stdout
+
+let approx ?(tol = 1e-9) a b = abs_float (a -. b) < tol
+
+let nt s = Grammar.Nonterm s
+let tm s = Grammar.Term s
+let prod lhs rhs = { Grammar.lhs; rhs; count = 1.0 }
+let g start ps = Grammar.make ~start ps
+
+let has_prod g lhs rhs =
+ List.exists
+ (fun (p : Grammar.production) ->
+ String.equal p.lhs lhs && Grammar.equal_rhs p.rhs rhs)
+ g.Grammar.productions
+
+let sample tokens count = { Corpus.tokens; count }
+
+let () =
+ check "log_gamma(1)=0" (approx (Scoring.log_gamma 1.0) 0.0);
+ check "log_gamma(2)=0" (approx (Scoring.log_gamma 2.0) 0.0);
+ check "log_gamma(1/2)=log sqrt pi"
+ (approx (Scoring.log_gamma 0.5) (0.5 *. log Float.pi));
+ check "log_gamma(3)=log 2"
+ (approx (Scoring.log_gamma 3.0) (log 2.0))
+
+let () =
+ let amb = g "S" [ prod "S" [ nt "S"; nt "S" ]; prod "S" [ tm "a" ] ] in
+ let prep = Parse.prepare amb in
+ check "P(a) = 1/2" (approx (Parse.inside prep [ "a" ]) 0.5);
+ check "P(aa) = 1/8" (approx (Parse.inside prep [ "a"; "a" ]) 0.125);
+ check "P(aaa) = 1/16" (approx (Parse.inside prep [ "a"; "a"; "a" ]) 0.0625);
+ let vp, vt = Parse.viterbi prep [ "a"; "a" ] in
+ check "viterbi P(aa) = 1/8" (approx vp 0.125);
+ check "viterbi tree exists" (vt <> None)
+
+let () =
+ let gr = g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ] in
+ let prep = Parse.prepare gr in
+ let cnt, p = Parse.expected_counts prep [ "a"; "b" ] in
+ check "inside(ab)=1" (approx p 1.0);
+ check "E[count S->AB]=1" (approx cnt.(0) 1.0);
+ check "E[count A->a]=1" (approx cnt.(1) 1.0);
+ check "E[count B->b]=1" (approx cnt.(2) 1.0)
+
+let () =
+ let amb = g "S" [ prod "S" [ nt "S"; nt "S" ]; prod "S" [ tm "a" ] ] in
+ let prep = Parse.prepare amb in
+ let cnt, _ = Parse.expected_counts prep [ "a"; "a" ] in
+ check "E[count S->SS]=1" (approx cnt.(0) 1.0);
+ check "E[count S->a]=2" (approx cnt.(1) 2.0)
+
+let () =
+ let gd =
+ g "S"
+ [ prod "S" [ tm "a" ];
+ { Grammar.lhs = "S"; rhs = [ tm "a" ]; count = 2.0 } ]
+ in
+ let gn = Grammar.normalize gd in
+ check "dedupe keeps one production"
+ (List.length gn.Grammar.productions = 1);
+ check "dedupe sums counts"
+ (approx (List.hd gn.Grammar.productions).Grammar.count 3.0)
+
+let () =
+ let gs = g "S" [ prod "S" [ nt "S" ]; prod "S" [ tm "a" ] ] in
+ let gn = Grammar.normalize gs in
+ check "self-loop removed"
+ (not (has_prod gn "S" [ nt "S" ]));
+ check "self-loop grammar keeps S->a" (has_prod gn "S" [ tm "a" ])
+
+let () =
+ let gc =
+ Grammar.make ~start:"A"
+ [ prod "A" [ nt "B" ]; prod "B" [ nt "A" ]; prod "A" [ tm "a" ] ]
+ in
+ let gn = Grammar.normalize gc in
+ check "unit cycle removed" (not (has_prod gn "A" [ nt "B" ]));
+ check "unit cycle grammar keeps A->a" (has_prod gn "A" [ tm "a" ])
+
+let () =
+ let gu =
+ g "S" [ prod "S" [ nt "A" ]; prod "A" [ nt "B" ]; prod "B" [ tm "a" ] ]
+ in
+ let gn = Grammar.normalize gu in
+ check "unit productions are preserved" (has_prod gn "S" [ nt "A" ]);
+ check "unit productions chain preserved" (has_prod gn "A" [ nt "B" ])
+
+let () =
+ let cfg = { Scoring.default_config with alpha = 1.0; em_iters = 5 } in
+ let g1 = g "S" [ prod "S" [ tm "a" ] ] in
+ let m1, _, _, _ =
+ Scoring.marginal_loglik ~config:cfg g1 [ sample [ "a" ] 1.0 ]
+ in
+ check "marginal(S->a, {a}, alpha=1) = 0" (approx m1 0.0);
+ let g2 = g "S" [ prod "S" [ tm "a" ]; prod "S" [ tm "b" ] ] in
+ let m2, _, _, _ =
+ Scoring.marginal_loglik ~config:cfg g2 [ sample [ "a" ] 1.0 ]
+ in
+ check "marginal(S->a|b, {a}, alpha=1) = log 1/2"
+ (approx m2 (log 0.5))
+
+let paper_initial () =
+ g "S"
+ [
+ prod "S" [ nt "A"; nt "B" ];
+ prod "S" [ nt "A"; nt "A"; nt "B"; nt "B" ];
+ prod "S" [ nt "A"; nt "A"; nt "A"; nt "B"; nt "B"; nt "B" ];
+ prod "A" [ tm "a" ];
+ prod "B" [ tm "b" ];
+ ]
+
+let () =
+ let g0 = paper_initial () in
+ let gc1 = Grammar.normalize (Transform.apply g0 (Transform.Chunk [ nt "A"; nt "B" ])) in
+ check "chunk(AB): S->X" (has_prod gc1 "S" [ nt "X" ]);
+ check "chunk(AB): S->AXB" (has_prod gc1 "S" [ nt "A"; nt "X"; nt "B" ]);
+ check "chunk(AB): S->AAXBB"
+ (has_prod gc1 "S" [ nt "A"; nt "A"; nt "X"; nt "B"; nt "B" ]);
+ check "chunk(AB): X->AB" (has_prod gc1 "X" [ nt "A"; nt "B" ]);
+ check "chunk preserves language"
+ (Grammar.language_up_to g0 6 = Grammar.language_up_to gc1 6)
+
+let () =
+ let g0 = paper_initial () in
+ let gc1 = Grammar.normalize (Transform.apply g0 (Transform.Chunk [ nt "A"; nt "B" ])) in
+ let gc2 =
+ Grammar.normalize (Transform.apply gc1 (Transform.Chunk [ nt "A"; nt "X"; nt "B" ]))
+ in
+ check "chunk(AXB): S->Y" (has_prod gc2 "S" [ nt "Y" ]);
+ check "chunk(AXB): S->AYB" (has_prod gc2 "S" [ nt "A"; nt "Y"; nt "B" ]);
+ check "chunk(AXB): Y->AXB" (has_prod gc2 "Y" [ nt "A"; nt "X"; nt "B" ]);
+ let gm = Grammar.normalize (Transform.apply gc2 (Transform.Merge ("S", "Y"))) in
+ let gf = Grammar.normalize (Transform.apply gm (Transform.Merge ("S", "X"))) in
+ check "final: S->AB" (has_prod gf "S" [ nt "A"; nt "B" ]);
+ check "final: S->ASB" (has_prod gf "S" [ nt "A"; nt "S"; nt "B" ]);
+ let expected =
+ [
+ [ "a"; "b" ];
+ [ "a"; "a"; "b"; "b" ];
+ [ "a"; "a"; "a"; "b"; "b"; "b" ];
+ [ "a"; "a"; "a"; "a"; "b"; "b"; "b"; "b" ];
+ ]
+ |> List.sort compare
+ in
+ check "final grammar generates a^n b^n"
+ (Grammar.language_up_to gf 8 = expected)
+
+let () =
+ let g0 =
+ g "S"
+ [ prod "S" [ nt "A" ]; prod "S" [ nt "B" ]; prod "A" [ tm "a" ];
+ prod "B" [ tm "b" ] ]
+ in
+ let gm = Grammar.normalize (Transform.apply g0 (Transform.Merge ("A", "B"))) in
+ check "merge unifies terminals into one nonterminal"
+ (List.length (Grammar.terminals gm) = 2
+ && List.length
+ (List.filter (fun p -> Grammar.symbol_is_nonterm (List.hd p.Grammar.rhs))
+ gm.Grammar.productions)
+ >= 1);
+ check "merge keeps both terminal productions" (has_prod gm "A" [ tm "a" ] && has_prod gm "A" [ tm "b" ])
+
+let () =
+ let cfg =
+ {
+ Scoring.default_config with
+ em_iters = 5;
+ max_chunk = 3;
+ prior_weight = 1.0;
+ }
+ in
+ let corpus = [ sample [ "a"; "b" ] 1.0; sample [ "a"; "a"; "b"; "b" ] 1.0; sample [ "a"; "a"; "a"; "b"; "b"; "b" ] 1.0 ] in
+ let result =
+ Search.run ~config:cfg ~beam_width:3 ~max_steps:12 ~patience:3
+ (paper_initial ()) corpus
+ in
+ check "search improves the posterior"
+ (result.Search.best_score > result.Search.initial_score);
+ check "search returns accepted steps"
+ (List.length result.Search.steps > 0);
+ check "search output generates training strings"
+ (List.for_all
+ (fun s ->
+ List.mem s (Grammar.language_up_to result.Search.best 8))
+ [ [ "a"; "b" ]; [ "a"; "a"; "b"; "b" ]; [ "a"; "a"; "a"; "b"; "b"; "b" ] ])
+
+let () =
+ let g0 =
+ g "S"
+ [ prod "S" [ nt "A"; nt "S"; nt "A" ]; prod "S" [ tm "c" ];
+ prod "S" [ nt "A"; nt "A"; nt "S"; nt "A"; nt "A" ];
+ prod "A" [ tm "a" ] ]
+ in
+ let gn = Grammar.normalize g0 in
+ check "redundant production pruned"
+ (not (has_prod gn "S" [ nt "A"; nt "A"; nt "S"; nt "A"; nt "A" ]));
+ check "recursive production kept" (has_prod gn "S" [ nt "A"; nt "S"; nt "A" ]);
+ check "pruning preserves the language"
+ (Grammar.language_up_to g0 5 = Grammar.language_up_to gn 5)
+
+let () =
+ let gr =
+ g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ]
+ in
+ let prep = Parse.prepare gr in
+ let _, tree = Parse.viterbi prep [ "a"; "b" ] in
+ match tree with
+ | Some t ->
+ check "parse tree yield" (Parse.tree_yield t = [ "a"; "b" ]);
+ check "parse tree root" (match t with Parse.Node ("S", _) -> true | _ -> false)
+ | None -> check "parse tree yield" false
+
+let () =
+ let gr =
+ g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ]
+ in
+ let corpus = [ sample [ "a"; "b" ] 1.0 ] in
+ let avg, _, _, unp = Scoring.heldout gr ~train:corpus corpus in
+ check "heldout loglik per token = 0" (approx avg 0.0);
+ check "heldout parseable" (unp = 0)
+
+let contains hay needle =
+ let n = String.length hay and m = String.length needle in
+ let rec go i = if i + m > n then false else if String.sub hay i m = needle then true else go (i + 1) in
+ m = 0 || go 0
+
+let () =
+ let g0 = g "S" [ prod "S" [ nt "S"; nt "S" ]; prod "S" [ tm "a" ] ] in
+ let d = Dot.dot_grammar g0 in
+ check "dot numbers repeated rhs symbols in order"
+ (contains d "pr_0 -> nt_S [label=\"1\"];"
+ && contains d "pr_0 -> nt_S [label=\"2\"];");
+ check "dot labels production probability" (contains d "p=0.500");
+ check "dot marks the start symbol" (contains d "shape=doublecircle");
+ let tree_dot = Dot.dot_tree (Parse.Node ("S", [ Parse.Leaf "a" ])) in
+ check "parse tree dot has an edge" (contains tree_dot "->")
+
+let () =
+ check "digamma(1)" (approx (Scoring.digamma 1.0) (-0.5772156649015329) ~tol:1e-9);
+ check "digamma(2)" (approx (Scoring.digamma 2.0) 0.42278433509846713 ~tol:1e-9)
+
+let () =
+ let mk mode =
+ { Scoring.default_config with alpha = 1.0; em_iters = 10; mode }
+ in
+ let corpus = [ sample [ "a" ] 1.0 ] in
+ let g1 = g "S" [ prod "S" [ tm "a" ] ] in
+ let d1 = Scoring.details ~config:(mk Scoring.Variational) g1 corpus in
+ check "variational ELBO (single production) = 0"
+ (approx d1.Scoring.posterior 0.0 ~tol:1e-6);
+ let g2 = g "S" [ prod "S" [ tm "a" ]; prod "S" [ tm "b" ] ] in
+ let dm = Scoring.details ~config:(mk Scoring.Marginal) g2 corpus in
+ let dv = Scoring.details ~config:(mk Scoring.Variational) g2 corpus in
+ check "variational matches marginal (unambiguous)"
+ (approx dv.Scoring.posterior dm.Scoring.posterior ~tol:1e-6)
+
+let () =
+ let g0 =
+ g "S"
+ [ prod "S" [ nt "A"; nt "B"; nt "A"; nt "B" ]; prod "A" [ tm "a" ];
+ prod "B" [ tm "b" ] ]
+ in
+ let all =
+ Grammar.normalize (Transform.apply g0 (Transform.Chunk [ nt "A"; nt "B" ]))
+ in
+ let first =
+ Grammar.normalize
+ (Transform.apply ~occurrence:Scoring.First_occurrence g0
+ (Transform.Chunk [ nt "A"; nt "B" ]))
+ in
+ check "chunk all replaces every occurrence"
+ (has_prod all "S" [ nt "X"; nt "X" ]);
+ check "chunk first replaces only the leftmost"
+ (has_prod first "S" [ nt "X"; nt "A"; nt "B" ])
+
+let () =
+ let g0 =
+ g "S"
+ [ prod "S" [ nt "A"; nt "S"; nt "A" ]; prod "S" [ tm "c" ];
+ prod "S" [ nt "A"; nt "A"; nt "A"; nt "S"; nt "A"; nt "A"; nt "A" ];
+ prod "A" [ tm "a" ] ]
+ in
+ let gn = Grammar.normalize g0 in
+ check "deeply redundant production pruned"
+ (not
+ (has_prod gn "S"
+ [ nt "A"; nt "A"; nt "A"; nt "S"; nt "A"; nt "A"; nt "A" ]));
+ check "deep redundancy preserves language"
+ (Grammar.language_up_to g0 7 = Grammar.language_up_to gn 7)
+
+let () =
+ let g0 =
+ g "S"
+ [ prod "S" [ nt "A"; nt "B" ]; prod "S" [ nt "A"; nt "S"; nt "B" ];
+ prod "S" [ nt "A"; nt "A"; nt "B" ]; prod "A" [ tm "a" ];
+ prod "B" [ tm "b" ] ]
+ in
+ let gn = Grammar.normalize g0 in
+ check "non-redundant production kept"
+ (has_prod gn "S" [ nt "A"; nt "A"; nt "B" ])
+
+let () =
+ if !failures = 0 then Printf.printf "\nAll tests passed.\n"
+ else begin
+ Printf.printf "\n%d test(s) failed.\n" !failures;
+ exit 1
+ end