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

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)