This commit is contained in:
sneeker committed 2026-09-18 16:27:30 +00:00
commit 100e5a8239
96 files changed
+39927

No files matched your search

+67
View File
@@ -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
View File
@@ -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
+4
View File
@@ -0,0 +1,4 @@
(library
(name scfg)
(flags (:standard -warn-error -a))
(modules grammar corpus parse scoring transform search dot experiment))
+282
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+193
View File
@@ -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;
}