From 84f6a9fdb14cab85e666b94822800ea96d44aed2 Mon Sep 17 00:00:00 2001 From: sneeker Date: Sat, 29 Apr 2017 16:18:00 +0000 Subject: [PATCH] Eliminate empty changes and unused cache slots --- src/incremental.ml | 23 ++++--- src/main.ml | 9 ++- src/simplify.ml | 132 +++++++++++++++++++++++++++++++++++++++ src/simplify.mli | 3 + test/test_incremental.ml | 80 ++++++++++++++++++++++++ test/test_main.ml | 1 + 6 files changed, 236 insertions(+), 12 deletions(-) create mode 100644 src/simplify.ml create mode 100644 src/simplify.mli diff --git a/src/incremental.ml b/src/incremental.ml index c80fa4d..ca53ded 100644 --- a/src/incremental.ml +++ b/src/incremental.ml @@ -230,17 +230,20 @@ let rec propagate context state key up_old up_new node_id = | None -> Delta_runtime.count_mapping context.ctx_counters; Some (eval_scalar [ (Ident.stamp parameter, value) ] body) - | Some previous -> - Delta_runtime.count_scalar_delta context.ctx_counters; - let delta = - delta_scalar plan - [ (Ident.stamp parameter, previous) ] - [ (Ident.stamp parameter, value) ] - body - in + | Some previous -> ( match old_value with - | Some cached_value -> Some (Change.apply cached_value delta) - | None -> Some (eval_scalar [ (Ident.stamp parameter, value) ] body)) + | Some cached_value -> + Delta_runtime.count_scalar_delta context.ctx_counters; + let delta = + delta_scalar plan + [ (Ident.stamp parameter, previous) ] + [ (Ident.stamp parameter, value) ] + body + in + Some (Change.apply cached_value delta) + | None -> + Delta_runtime.count_mapping context.ctx_counters; + Some (eval_scalar [ (Ident.stamp parameter, value) ] body))) in let cached = match (old_value, new_value) with diff --git a/src/main.ml b/src/main.ml index e5c306f..d5ca9ca 100644 --- a/src/main.ml +++ b/src/main.ml @@ -95,7 +95,9 @@ let specialize_source path = Specialize.program (infer_source path) let anf_source path = Anf.program (specialize_source path) -let plan_source path = Graph.build (anf_source path) +let graph_source path = Graph.build (anf_source path) + +let plan_source path = Simplify.simplify (graph_source path) let frontend_unavailable () = Diagnostic.error Location.none "the delta pipeline beyond parsing is not implemented in this revision" @@ -106,7 +108,10 @@ let check path = let dump stage path = match stage with | "anf" -> print_string (Anf.program_to_string (anf_source path)) - | "delta" -> print_string (Incremental.dump_delta (plan_source path)) + | "delta" -> + let graph = graph_source path in + print_string + (Simplify.decisions graph ^ Incremental.dump_delta (Simplify.simplify graph)) | _ -> let program = infer_source path in match stage with diff --git a/src/simplify.ml b/src/simplify.ml new file mode 100644 index 0000000..b98a997 --- /dev/null +++ b/src/simplify.ml @@ -0,0 +1,132 @@ +type action = + | Keep of Graph.cache + | Elide + +let rec uses_division expr = + match expr.Anf.a with + | Anf.ABinop (Syntax.Div, _, _) -> true + | Anf.ABinop (_, _, _) -> false + | Anf.AIf (_, then_branch, else_branch) -> + uses_division then_branch || uses_division else_branch + | Anf.ALet (_, bound, body) -> uses_division bound || uses_division body + | Anf.AAtom _ | Anf.ATuple _ | Anf.ARecord _ | Anf.AField _ | Anf.AApp _ | Anf.ALambda _ -> false + | Anf.AFilter _ | Anf.AMap _ | Anf.ASum _ | Anf.ACount _ -> false + +let identity_projection parameter body = + match body.Anf.a with Anf.AAtom (Anf.AVar ident) -> Ident.equal ident parameter | _ -> false + +let constant_true body = match body.Anf.a with Anf.AAtom (Anf.ABool true) -> true | _ -> false + +let count_only plan node_id = + let rec reaches_count seen node_id = + if Util.contains node_id seen then false + else + let node = Graph.node_of_id plan node_id in + match node.Graph.n_kind with + | Graph.Count -> true + | Graph.Filter _ -> List.for_all (reaches_count (node_id :: seen)) plan.Graph.pl_consumers.(node_id) + | Graph.Map _ -> List.for_all (reaches_count (node_id :: seen)) plan.Graph.pl_consumers.(node_id) + | Graph.Source | Graph.Sum -> false + in + plan.Graph.pl_consumers.(node_id) <> [] + && List.for_all (reaches_count [ node_id ]) plan.Graph.pl_consumers.(node_id) + +let root_id plan = match plan.Graph.pl_result with Graph.Result_collection id -> Some id | _ -> None + +let decide plan node = + let is_root = match root_id plan with Some id -> id = node.Graph.n_id | None -> false in + match node.Graph.n_kind with + | Graph.Source | Graph.Sum | Graph.Count -> Keep node.Graph.n_cache + | Graph.Filter (_, body) -> if constant_true body then Elide else Keep node.Graph.n_cache + | Graph.Map (parameter, body) -> + if identity_projection parameter body then Elide + else if is_root then Keep node.Graph.n_cache + else if count_only plan node.Graph.n_id then + if uses_division body then Keep Graph.No_cache else Elide + else Keep node.Graph.n_cache + +let simplify plan = + let actions = List.map (fun node -> (node.Graph.n_id, decide plan node)) plan.Graph.pl_nodes in + let action_of id = Util.assoc_opt id actions in + let resolved = Hashtbl.create 16 in + let rec resolve id = + match Util.hashtbl_find_opt resolved id with + | Some id -> id + | None -> + let node = Graph.node_of_id plan id in + let result = + match (action_of id, node.Graph.n_input) with + | Some Elide, Some input -> resolve input + | Some Elide, None -> id + | _ -> + let id = node.Graph.n_id in + Hashtbl.replace resolved id id; + id + in + Hashtbl.replace resolved id result; + result + in + let kept = + List.filter + (fun node -> + match action_of node.Graph.n_id with + | Some (Keep _) -> true + | Some Elide -> false + | None -> false) + plan.Graph.pl_nodes + in + let renumbered = List.mapi (fun index node -> (node.Graph.n_id, index)) kept in + let new_id old_id = + match Util.assoc_opt (resolve old_id) renumbered with + | Some id -> id + | None -> -1 + in + let nodes = + List.map + (fun (old_id, fresh_id) -> + let node = Graph.node_of_id plan old_id in + let cache = + match action_of old_id with Some (Keep cache) -> cache | _ -> Graph.No_cache + in + let input = match node.Graph.n_input with None -> None | Some input -> Some (new_id input) in + { node with Graph.n_id = fresh_id; n_input = input; n_cache = cache }) + renumbered + in + let node_count = List.length nodes in + let consumers = Array.make node_count [] in + List.iter + (fun node -> + match node.Graph.n_input with + | Some input when input >= 0 && input < node_count -> + consumers.(input) <- consumers.(input) @ [ node.Graph.n_id ] + | _ -> ()) + nodes; + let result = + match plan.Graph.pl_result with + | Graph.Result_collection id -> Graph.Result_collection (new_id id) + | Graph.Result_scalar expr -> Graph.Result_scalar expr + in + let scalar_bindings = + List.map (fun (stamp, id) -> (stamp, new_id id)) plan.Graph.pl_scalar_bindings + in + { + plan with + Graph.pl_nodes = nodes; + pl_result = result; + pl_consumers = consumers; + pl_scalar_bindings = scalar_bindings; + } + +let decisions plan = + let lines = + List.map + (fun node -> + match decide plan node with + | Keep cache -> + Printf.sprintf " keep node %d: %s [cache: %s]" node.Graph.n_id + (Graph.kind_to_string node) (Graph.cache_to_string cache) + | Elide -> + Printf.sprintf " elide node %d: %s" node.Graph.n_id (Graph.kind_to_string node)) + plan.Graph.pl_nodes + in + Util.join "\n" lines ^ "\n" diff --git a/src/simplify.mli b/src/simplify.mli new file mode 100644 index 0000000..8fa69fc --- /dev/null +++ b/src/simplify.mli @@ -0,0 +1,3 @@ +val simplify : Graph.plan -> Graph.plan +val decisions : Graph.plan -> string +val uses_division : Anf.expr -> bool diff --git a/test/test_incremental.ml b/test/test_incremental.ml index e7e5a7d..0ca958f 100644 --- a/test/test_incremental.ml +++ b/test/test_incremental.ml @@ -661,3 +661,83 @@ let aggregate_cases = fail "cached matches the reference" (Printf.sprintf "ops=%d" (List.length ops))) batches ); ] + +let simplified text = Simplify.simplify (plan text) + +let simplify_cases = + [ + ( "a map in front of count is elided when the projection is total", + fun () -> + let plan = simplified "input rows : collection int\nquery q = rows |> map (fun r -> r * 2) |> count\n" in + check_equal_int "two nodes" 2 (List.length plan.Graph.pl_nodes); + check "source then count" + (match (List.nth plan.Graph.pl_nodes 1).Graph.n_kind with Graph.Count -> true | _ -> false) ); + ( "a map that can divide keeps its node but loses its cache", + fun () -> + let graph = + Graph.build + (Anf.program + (Specialize.program + (infer "input rows : collection int\nquery q = rows |> map (fun r -> 100 / r) |> count\n"))) + in + let plan = Simplify.simplify graph in + check_equal_int "three nodes" 3 (List.length plan.Graph.pl_nodes); + let mapper = List.nth plan.Graph.pl_nodes 1 in + check "kept as a map" (match mapper.Graph.n_kind with Graph.Map _ -> true | _ -> false); + check "no cache" (mapper.Graph.n_cache = Graph.No_cache) ); + ( "a map feeding sum keeps its cache", + fun () -> + let plan = simplified "input rows : collection int\nquery q = rows |> map (fun r -> r * 2) |> sum\n" in + check_equal_int "three nodes" 3 (List.length plan.Graph.pl_nodes); + check "cache kept" ((List.nth plan.Graph.pl_nodes 1).Graph.n_cache = Graph.Cached_values) ); + ( "an identity map is elided", + fun () -> + let plan = simplified "input rows : collection int\nquery q = rows |> map (fun r -> r)\n" in + check_equal_int "one node" 1 (List.length plan.Graph.pl_nodes); + check "source only" (match (List.nth plan.Graph.pl_nodes 0).Graph.n_kind with Graph.Source -> true | _ -> false) ); + ( "a filter with a constant true predicate is elided", + fun () -> + let plan = simplified "input rows : collection int\nquery q = rows |> filter (fun r -> true) |> sum\n" in + check_equal_int "two nodes" 2 (List.length plan.Graph.pl_nodes); + check "a sum remains" (match (List.nth plan.Graph.pl_nodes 1).Graph.n_kind with Graph.Sum -> true | _ -> false) ); + ( "a filter is never elided when its predicate depends on the row", + fun () -> + let plan = simplified "input rows : collection int\nquery q = rows |> filter (fun r -> r > 0) |> count\n" in + check_equal_int "three nodes" 3 (List.length plan.Graph.pl_nodes); + check "filter kept" (match (List.nth plan.Graph.pl_nodes 1).Graph.n_kind with Graph.Filter _ -> true | _ -> false) ); + ( "the simplified plan still computes reference results", + fun () -> + let f = + build_fixture "input rows : collection int\nquery q = rows |> map (fun r -> 100 / r) |> sum\n" + [ (1, Value.VInt 5); (2, Value.VInt 4) ] + in + check_equal_string "initial" (Value.to_string (reference_result f f.fx_state)) + (Value.to_string (Incremental.result f.fx_plan f.fx_state)); + (match step f [ Change.OpInsert (3, Value.VInt 10) ] with + | Error message -> fail "insert" message + | Ok (applied, cached, _) -> + check_equal_string "applied" (Value.to_string (reference_result f f.fx_state)) + (Value.to_string applied); + check_equal_string "cached" (Value.to_string (reference_result f f.fx_state)) + (Value.to_string cached)) ); + ( "an elided map still reports division errors for new rows", + fun () -> + let f = + build_fixture "input rows : collection int\nquery q = rows |> map (fun r -> 100 / r) |> count\n" + [ (1, Value.VInt 5) ] + in + (match step f [ Change.OpInsert (2, Value.VInt 0) ] with + | Ok _ -> fail "division" "expected the elided map to still evaluate" + | Error message -> check "division by zero" (String.length message > 0)) ); + ( "decisions are reported per node", + fun () -> + let graph = + Graph.build + (Anf.program + (Specialize.program + (infer "input rows : collection int\nquery q = rows |> map (fun r -> r) |> count\n"))) + in + let report = Simplify.decisions graph in + check "the identity map is elided" (String.length report > 0); + check "mentions a source" (Util.starts_with " keep node 0: source" report) ); + ] diff --git a/test/test_main.ml b/test/test_main.ml index 7d9aab3..88148da 100644 --- a/test/test_main.ml +++ b/test/test_main.ml @@ -383,5 +383,6 @@ let () = Test_harness.run_suite "executor" Test_incremental.executor_cases; Test_harness.run_suite "filter" Test_incremental.filter_cases; Test_harness.run_suite "aggregates" Test_incremental.aggregate_cases; + Test_harness.run_suite "simplify" Test_incremental.simplify_cases; Printf.printf "%d cases, %d failures\n" (Test_harness.case_count ()) (Test_harness.failure_count ()); exit (if Test_harness.failure_count () = 0 then 0 else 1)