initial
This commit is contained in:
commit
3170b7b398
96 files changed
+39927
No files matched your search
@@ -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
|
||||
+143
@@ -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
|
||||
@@ -0,0 +1,4 @@
|
||||
(library
|
||||
(name scfg)
|
||||
(flags (:standard -warn-error -a))
|
||||
(modules grammar corpus parse scoring transform search dot experiment))
|
||||
@@ -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
|
||||
+522
@@ -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)
|
||||
+355
@@ -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)
|
||||
+278
@@ -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)
|
||||
+250
@@ -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)
|
||||
@@ -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;
|
||||
}
|
||||
Reference in new issue
Block a user