fix(scoring): distinguish approximate evidence from exact marginal likelihood

This commit is contained in:
milner committed 2026-09-23 13:10:00 +00:00
1 parent 697801a0ee
commit a4db48b636
4 files changed
+30 -22

No files matched your search

+1 -1
View File
@@ -9,6 +9,6 @@ We start from the most specific grammar the data permits: every sample contribut
From there, we generalise with two operators, merging and chunking. Merging takes a pair of nonterminals and folds them into a single nonterminal containing the union of their productions; chunking replaces a contiguous sequence of symbols with a fresh nonterminal. Chunking doesn't itself change the language the grammar generates, but it changes the internal structure in a way that can expose useful merges which weren't previously available. From there, we generalise with two operators, merging and chunking. Merging takes a pair of nonterminals and folds them into a single nonterminal containing the union of their productions; chunking replaces a contiguous sequence of symbols with a fresh nonterminal. Chunking doesn't itself change the language the grammar generates, but it changes the internal structure in a way that can expose useful merges which weren't previously available.
We rank candidate grammars by the posterior `P(M | X) ∝ P(M) P(X | M)`, where the prior is a description length and the likelihood integrates over the production probabilities under symmetric Dirichlet priors. The scoring therefore accounts for uncertainty in the production probabilities, rather than relying on a single fitted parameterisation. We rank candidate grammars with a structural description length and an approximate evidence objective. The default scorer fits production probabilities with expectation maximisation, computes fractional expected production counts, and applies a symmetric Dirichlet correction. Maximum likelihood and a variational evidence lower bound are also available.
We explore the resulting grammar space using either beam or best first search, then fit the parameters by expectation maximisation once the grammar structure is fixed. Parsing is a generalised CYK or inside computation over spans, which lets the grammar be scored directly without an intermediate conversion into Chomsky normal form. We explore the resulting grammar space using either beam or best first search, then fit the parameters by expectation maximisation once the grammar structure is fixed. Parsing is a generalised CYK or inside computation over spans, which lets the grammar be scored directly without an intermediate conversion into Chomsky normal form.
+9 -6
View File
@@ -46,10 +46,11 @@ let config_of_opts opts =
em_iters = int_opt opts "em-iters" Scoring.default_config.em_iters; em_iters = int_opt opts "em-iters" Scoring.default_config.em_iters;
max_chunk = int_opt opts "max-chunk" Scoring.default_config.max_chunk; max_chunk = int_opt opts "max-chunk" Scoring.default_config.max_chunk;
mode = mode =
(match opt opts "mode" "marginal" with (match opt opts "mode" "approximate" with
| "ml" -> Scoring.Maximum_likelihood | "ml" -> Scoring.Maximum_likelihood
| "variational" -> Scoring.Variational | "variational" -> Scoring.Variational
| _ -> Scoring.Marginal); | "approximate" | "marginal" -> Scoring.Approximate_evidence
| mode -> invalid_arg ("unknown scoring mode: " ^ mode));
chunk_occurrence = chunk_occurrence =
(if opt opts "chunk-occurrence" "all" = "first" then (if opt opts "chunk-occurrence" "all" = "first" then
Scoring.First_occurrence Scoring.First_occurrence
@@ -63,7 +64,7 @@ let usage () =
Commands:\n\ Commands:\n\
\ induce --corpus FILE [--out DIR] [--beam N] [--max-steps N]\n\ \ induce --corpus FILE [--out DIR] [--beam N] [--max-steps N]\n\
\ [--patience N] [--alpha A] [--prior W] [--max-chunk N]\n\ \ [--patience N] [--alpha A] [--prior W] [--max-chunk N]\n\
\ [--em-iters N] [--mode marginal|ml|variational]\n\ \ [--em-iters N] [--mode approximate|ml|variational]\n\
\ [--chunk-occurrence all|first] [--render]\n\ \ [--chunk-occurrence all|first] [--render]\n\
\ [--search beam|bestfirst] [--frontier N] [--max-expansions N]\n\ \ [--search beam|bestfirst] [--frontier N] [--max-expansions N]\n\
\ score --grammar FILE --corpus FILE [--alpha A] [--prior W]\n\ \ score --grammar FILE --corpus FILE [--alpha A] [--prior W]\n\
@@ -159,8 +160,8 @@ let cmd_induce opts =
result.Search.iterations (List.length result.Search.steps); result.Search.iterations (List.length result.Search.steps);
Printf.printf "posterior: %.4f -> %.4f\n" result.Search.initial_score Printf.printf "posterior: %.4f -> %.4f\n" result.Search.initial_score
result.Search.best_score; result.Search.best_score;
Printf.printf "structure prior %.4f, marginal loglik %.4f, dl %.1f bits\n" Printf.printf "structure prior %.4f, objective %.4f, dl %.1f bits\n"
details.Scoring.prior details.Scoring.marginal details.Scoring.dl_bits; details.Scoring.prior details.Scoring.objective details.Scoring.dl_bits;
Printf.printf "final grammar: %d nonterminals, %d productions\n" Printf.printf "final grammar: %d nonterminals, %d productions\n"
details.Scoring.num_nonterminals details.Scoring.num_productions; details.Scoring.num_nonterminals details.Scoring.num_productions;
print_endline "\nfinal grammar:"; print_endline "\nfinal grammar:";
@@ -178,7 +179,9 @@ let cmd_score opts =
let d = Scoring.details ~config g corpus in let d = Scoring.details ~config g corpus in
Printf.printf "posterior %.6f\n" d.Scoring.posterior; Printf.printf "posterior %.6f\n" d.Scoring.posterior;
Printf.printf "structural prior %.6f\n" d.Scoring.prior; Printf.printf "structural prior %.6f\n" d.Scoring.prior;
Printf.printf "marginal loglik %.6f\n" d.Scoring.marginal; Printf.printf "objective %.6f\n" d.Scoring.objective;
Printf.printf "approximate evidence %.6f\n"
d.Scoring.approximate_evidence;
Printf.printf "ML corpus loglik %.6f\n" d.Scoring.ml_loglik; Printf.printf "ML corpus loglik %.6f\n" d.Scoring.ml_loglik;
Printf.printf "description length %.2f bits\n" d.Scoring.dl_bits; Printf.printf "description length %.2f bits\n" d.Scoring.dl_bits;
Printf.printf "nonterminals %d\n" d.Scoring.num_nonterminals; Printf.printf "nonterminals %d\n" d.Scoring.num_nonterminals;
+12 -9
View File
@@ -1,4 +1,4 @@
type score_mode = Marginal | Maximum_likelihood | Variational type score_mode = Approximate_evidence | Maximum_likelihood | Variational
type chunk_occurrence = All_occurrences | First_occurrence type chunk_occurrence = All_occurrences | First_occurrence
@@ -19,7 +19,7 @@ let default_config =
em_iters = 8; em_iters = 8;
em_tol = 1e-7; em_tol = 1e-7;
max_chunk = 2; max_chunk = 2;
mode = Marginal; mode = Approximate_evidence;
chunk_occurrence = All_occurrences; chunk_occurrence = All_occurrences;
} }
@@ -95,7 +95,8 @@ let description_length_bits (g : Grammar.t) =
let structural_logprior ?(config = default_config) (g : Grammar.t) = let structural_logprior ?(config = default_config) (g : Grammar.t) =
-.config.prior_weight *. description_length_bits g *. log 2.0 -.config.prior_weight *. description_length_bits g *. log 2.0
let marginal_loglik ?(config = default_config) (g : Grammar.t) (corpus : Corpus.t) = let approximate_log_evidence ?(config = default_config) (g : Grammar.t)
(corpus : Corpus.t) =
let uprep = Parse.prepare_uniform g in let uprep = Parse.prepare_uniform g in
if if
@@ -228,15 +229,16 @@ let posterior ?(config = default_config) g corpus =
match config.mode with match config.mode with
| Variational -> prior +. variational_loglik ~config g corpus | Variational -> prior +. variational_loglik ~config g corpus
| _ -> | _ ->
let marg, ll, _, _ = marginal_loglik ~config g corpus in let evidence, ll, _, _ = approximate_log_evidence ~config g corpus in
(match config.mode with (match config.mode with
| Marginal -> prior +. marg | Approximate_evidence -> prior +. evidence
| Maximum_likelihood -> prior +. ll | Maximum_likelihood -> prior +. ll
| Variational -> assert false) | Variational -> assert false)
type details = { type details = {
prior : float; prior : float;
marginal : float; objective : float;
approximate_evidence : float;
ml_loglik : float; ml_loglik : float;
posterior : float; posterior : float;
dl_bits : float; dl_bits : float;
@@ -245,18 +247,19 @@ type details = {
} }
let details ?(config = default_config) (g : Grammar.t) (corpus : Corpus.t) = let details ?(config = default_config) (g : Grammar.t) (corpus : Corpus.t) =
let marg, ll, _, _ = marginal_loglik ~config g corpus in let evidence, ll, _, _ = approximate_log_evidence ~config g corpus in
let prior = structural_logprior ~config g in let prior = structural_logprior ~config g in
let dl = description_length_bits g in let dl = description_length_bits g in
let likelihood = let likelihood =
match config.mode with match config.mode with
| Marginal -> marg | Approximate_evidence -> evidence
| Maximum_likelihood -> ll | Maximum_likelihood -> ll
| Variational -> variational_loglik ~config g corpus | Variational -> variational_loglik ~config g corpus
in in
{ {
prior; prior;
marginal = likelihood; objective = likelihood;
approximate_evidence = evidence;
ml_loglik = ll; ml_loglik = ll;
posterior = prior +. likelihood; posterior = prior +. likelihood;
dl_bits = dl; dl_bits = dl;
+8 -6
View File
@@ -148,14 +148,14 @@ let () =
let cfg = { Scoring.default_config with alpha = 1.0; em_iters = 5 } in let cfg = { Scoring.default_config with alpha = 1.0; em_iters = 5 } in
let g1 = g "S" [ prod "S" [ tm "a" ] ] in let g1 = g "S" [ prod "S" [ tm "a" ] ] in
let m1, _, _, _ = let m1, _, _, _ =
Scoring.marginal_loglik ~config:cfg g1 [ sample [ "a" ] 1.0 ] Scoring.approximate_log_evidence ~config:cfg g1 [ sample [ "a" ] 1.0 ]
in in
check "marginal(S->a, {a}, alpha=1) = 0" (approx m1 0.0); check "evidence(S->a, {a}, alpha=1) = 0" (approx m1 0.0);
let g2 = g "S" [ prod "S" [ tm "a" ]; prod "S" [ tm "b" ] ] in let g2 = g "S" [ prod "S" [ tm "a" ]; prod "S" [ tm "b" ] ] in
let m2, _, _, _ = let m2, _, _, _ =
Scoring.marginal_loglik ~config:cfg g2 [ sample [ "a" ] 1.0 ] Scoring.approximate_log_evidence ~config:cfg g2 [ sample [ "a" ] 1.0 ]
in in
check "marginal(S->a|b, {a}, alpha=1) = log 1/2" check "evidence(S->a|b, {a}, alpha=1) = log 1/2"
(approx m2 (log 0.5)) (approx m2 (log 0.5))
let paper_initial () = let paper_initial () =
@@ -308,9 +308,11 @@ let () =
check "variational ELBO (single production) = 0" check "variational ELBO (single production) = 0"
(approx d1.Scoring.posterior 0.0 ~tol:1e-6); (approx d1.Scoring.posterior 0.0 ~tol:1e-6);
let g2 = g "S" [ prod "S" [ tm "a" ]; prod "S" [ tm "b" ] ] in let g2 = g "S" [ prod "S" [ tm "a" ]; prod "S" [ tm "b" ] ] in
let dm = Scoring.details ~config:(mk Scoring.Marginal) g2 corpus in let dm =
Scoring.details ~config:(mk Scoring.Approximate_evidence) g2 corpus
in
let dv = Scoring.details ~config:(mk Scoring.Variational) g2 corpus in let dv = Scoring.details ~config:(mk Scoring.Variational) g2 corpus in
check "variational matches marginal (unambiguous)" check "variational matches approximate evidence (unambiguous)"
(approx dv.Scoring.posterior dm.Scoring.posterior ~tol:1e-6) (approx dv.Scoring.posterior dm.Scoring.posterior ~tol:1e-6)
let () = let () =