279 lines
8.2 KiB
OCaml
279 lines
8.2 KiB
OCaml
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)
|