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

+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)