Files
scfg-induction/lib/experiment.ml
T
2026-09-18 16:27:30 +00:00

283 lines
10 KiB
OCaml

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