initial
This commit is contained in:
commit
3170b7b398
96 files changed
+39927
No files matched your search
@@ -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
|
||||
Reference in new issue
Block a user