From 05769abc11f32d7838374b214cdf0672408e30ea Mon Sep 17 00:00:00 2001 From: sneeker Date: Sat, 25 Feb 2017 12:48:00 +0000 Subject: [PATCH] Specialize scalar helpers without variable capture --- src/main.ml | 8 +- src/specialize.ml | 222 +++++++++++++++++++++++++++++++++++++++++++++ src/specialize.mli | 1 + test/test_main.ml | 1 + test/test_type.ml | 115 +++++++++++++++++++++++ 5 files changed, 344 insertions(+), 3 deletions(-) create mode 100644 src/specialize.ml create mode 100644 src/specialize.mli diff --git a/src/main.ml b/src/main.ml index f6a00d5..aef6fb3 100644 --- a/src/main.ml +++ b/src/main.ml @@ -91,11 +91,13 @@ let resolve_source path = Resolve.program (parse_source path) let infer_source path = Infer.program (resolve_source path) +let specialize_source path = Specialize.program (infer_source path) + let frontend_unavailable () = Diagnostic.error Location.none "the delta pipeline beyond parsing is not implemented in this revision" let check path = - ignore (infer_source path) + ignore (specialize_source path) let dump stage path = let program = infer_source path in @@ -104,12 +106,12 @@ let dump stage path = | _ -> frontend_unavailable () let emit path output = - ignore (infer_source path); + ignore (specialize_source path); ignore output; frontend_unavailable () let build path output = - ignore (infer_source path); + ignore (specialize_source path); ignore output; frontend_unavailable () diff --git a/src/specialize.ml b/src/specialize.ml new file mode 100644 index 0000000..d606cb2 --- /dev/null +++ b/src/specialize.ml @@ -0,0 +1,222 @@ +type context = { + helpers : (int * Typed.expr) list; + node_limit : int; + mutable nodes : int; + mutable error : (Location.span * string) option; +} + +let make_context helpers node_limit = { helpers = helpers; node_limit = node_limit; nodes = 0; error = None } + +let count ctx span = + ctx.nodes <- ctx.nodes + 1; + if ctx.nodes > ctx.node_limit && ctx.error = None then + ctx.error <- + Some + ( span, + "specializing the helper functions expanded this query beyond the supported size" ) + +let ident table stamp = + let rec search = function + | [] -> None + | (candidate, replacement) :: rest -> if candidate = stamp then Some replacement else search rest + in + search table + +let rec subst table expr = + match expr.Typed.te with + | Typed.TVar ident' -> ( + match ident table (Ident.stamp ident') with + | Some replacement -> replacement + | None -> expr) + | Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TSource _ -> expr + | Typed.TLet (ident', bound, body) -> + let table = List.filter (fun (stamp, _) -> stamp <> Ident.stamp ident') table in + { expr with Typed.te = Typed.TLet (ident', subst table bound, subst table body) } + | Typed.TLambda (ident', body) -> + let table = List.filter (fun (stamp, _) -> stamp <> Ident.stamp ident') table in + { expr with Typed.te = Typed.TLambda (ident', subst table body) } + | Typed.TApp (fn, argument) -> + { expr with Typed.te = Typed.TApp (subst table fn, subst table argument) } + | Typed.TIf (condition, then_branch, else_branch) -> + { + expr with + Typed.te = Typed.TIf (subst table condition, subst table then_branch, subst table else_branch); + } + | Typed.TBinop (operator, left, right) -> + { expr with Typed.te = Typed.TBinop (operator, subst table left, subst table right) } + | Typed.TTuple items -> { expr with Typed.te = Typed.TTuple (List.map (subst table) items) } + | Typed.TRecord (name, fields) -> + { + expr with + Typed.te = Typed.TRecord (name, List.map (fun (label, value) -> (label, subst table value)) fields); + } + | Typed.TField (record, label) -> { expr with Typed.te = Typed.TField (subst table record, label) } + | Typed.TFilter (collection, predicate) -> + { expr with Typed.te = Typed.TFilter (subst table collection, subst table predicate) } + | Typed.TMap (collection, projection) -> + { expr with Typed.te = Typed.TMap (subst table collection, subst table projection) } + | Typed.TSum collection -> { expr with Typed.te = Typed.TSum (subst table collection) } + | Typed.TCount collection -> { expr with Typed.te = Typed.TCount (subst table collection) } + +let rec freshen table expr = + match expr.Typed.te with + | Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TSource _ -> expr + | Typed.TVar ident' -> ( + match ident table (Ident.stamp ident') with + | Some replacement -> { expr with Typed.te = Typed.TVar replacement } + | None -> expr) + | Typed.TLet (ident', bound, body) -> + let replacement = Ident.fresh (Ident.name ident') (Ident.span ident') in + { + expr with + Typed.te = + Typed.TLet + ( replacement, + freshen table bound, + freshen ((Ident.stamp ident', replacement) :: table) body ); + } + | Typed.TLambda (ident', body) -> + let replacement = Ident.fresh (Ident.name ident') (Ident.span ident') in + { + expr with + Typed.te = Typed.TLambda (replacement, freshen ((Ident.stamp ident', replacement) :: table) body); + } + | Typed.TApp (fn, argument) -> { expr with Typed.te = Typed.TApp (freshen table fn, freshen table argument) } + | Typed.TIf (condition, then_branch, else_branch) -> + { + expr with + Typed.te = + Typed.TIf (freshen table condition, freshen table then_branch, freshen table else_branch); + } + | Typed.TBinop (operator, left, right) -> + { expr with Typed.te = Typed.TBinop (operator, freshen table left, freshen table right) } + | Typed.TTuple items -> { expr with Typed.te = Typed.TTuple (List.map (freshen table) items) } + | Typed.TRecord (name, fields) -> + { + expr with + Typed.te = + Typed.TRecord (name, List.map (fun (label, value) -> (label, freshen table value)) fields); + } + | Typed.TField (record, label) -> { expr with Typed.te = Typed.TField (freshen table record, label) } + | Typed.TFilter (collection, predicate) -> + { expr with Typed.te = Typed.TFilter (freshen table collection, freshen table predicate) } + | Typed.TMap (collection, projection) -> + { expr with Typed.te = Typed.TMap (freshen table collection, freshen table projection) } + | Typed.TSum collection -> { expr with Typed.te = Typed.TSum (freshen table collection) } + | Typed.TCount collection -> { expr with Typed.te = Typed.TCount (freshen table collection) } + +let helper_of ctx ident = + let rec search = function + | [] -> None + | (stamp, body) :: rest -> if stamp = Ident.stamp ident then Some body else search rest + in + search ctx.helpers + +let is_collection expr = Types.is_collection expr.Typed.ty + +let rec specialize ctx expr = + count ctx expr.Typed.tspan; + match expr.Typed.te with + | Typed.TVar ident' -> ( + match helper_of ctx ident' with + | Some body -> specialize ctx (freshen [] body) + | None -> expr) + | Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TSource _ -> expr + | Typed.TLambda (ident', body) -> { expr with Typed.te = Typed.TLambda (ident', specialize ctx body) } + | Typed.TApp (fn, argument) -> + let fn = specialize ctx fn in + let argument = specialize ctx argument in + (match fn.Typed.te with + | Typed.TLambda (ident', body) -> + specialize ctx (subst [ (Ident.stamp ident', argument) ] body) + | _ -> { expr with Typed.te = Typed.TApp (fn, argument) }) + | Typed.TLet (ident', bound, body) -> + let bound = specialize ctx bound in + if is_collection bound then { expr with Typed.te = Typed.TLet (ident', bound, specialize ctx body) } + else specialize ctx (subst [ (Ident.stamp ident', bound) ] body) + | Typed.TIf (condition, then_branch, else_branch) -> + { + expr with + Typed.te = + Typed.TIf (specialize ctx condition, specialize ctx then_branch, specialize ctx else_branch); + } + | Typed.TBinop (operator, left, right) -> + { expr with Typed.te = Typed.TBinop (operator, specialize ctx left, specialize ctx right) } + | Typed.TTuple items -> { expr with Typed.te = Typed.TTuple (List.map (specialize ctx) items) } + | Typed.TRecord (name, fields) -> + { + expr with + Typed.te = + Typed.TRecord (name, List.map (fun (label, value) -> (label, specialize ctx value)) fields); + } + | Typed.TField (record, label) -> { expr with Typed.te = Typed.TField (specialize ctx record, label) } + | Typed.TFilter (collection, predicate) -> + let collection = specialize ctx collection in + let predicate = specialize ctx predicate in + check_function ctx predicate "predicate"; + { expr with Typed.te = Typed.TFilter (collection, predicate) } + | Typed.TMap (collection, projection) -> + let collection = specialize ctx collection in + let projection = specialize ctx projection in + check_function ctx projection "mapping function"; + { expr with Typed.te = Typed.TMap (collection, projection) } + | Typed.TSum collection -> { expr with Typed.te = Typed.TSum (specialize ctx collection) } + | Typed.TCount collection -> { expr with Typed.te = Typed.TCount (specialize ctx collection) } + +and check_function ctx expr what = + match expr.Typed.te with + | Typed.TLambda (parameter, body) -> + let free = free_variables [ Ident.stamp parameter ] body [] in + (match free with + | [] -> () + | (ident, span) :: _ -> + Diagnostic.error span + "the %s cannot be incrementalized because it uses `%s` from the surrounding scope" what + (Ident.display ident)) + | _ -> + Diagnostic.error expr.Typed.tspan + "the %s must be a function literal or a helper that specializes to one, but this expression is a function value" + what + +and free_variables bound expr acc = + match expr.Typed.te with + | Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit -> acc + | Typed.TSource ident -> (ident, expr.Typed.tspan) :: acc + | Typed.TVar ident -> + if List.exists (fun stamp -> stamp = Ident.stamp ident) bound then acc + else (ident, expr.Typed.tspan) :: acc + | Typed.TLet (ident, bound_expr, body) -> + let acc = free_variables bound bound_expr acc in + free_variables (Ident.stamp ident :: bound) body acc + | Typed.TLambda (ident, body) -> free_variables (Ident.stamp ident :: bound) body acc + | Typed.TApp (fn, argument) -> free_variables bound argument (free_variables bound fn acc) + | Typed.TIf (condition, then_branch, else_branch) -> + free_variables bound else_branch + (free_variables bound then_branch (free_variables bound condition acc)) + | Typed.TBinop (_, left, right) -> free_variables bound right (free_variables bound left acc) + | Typed.TTuple items -> List.fold_left (fun acc item -> free_variables bound item acc) acc items + | Typed.TRecord (_, fields) -> + List.fold_left (fun acc (_, value) -> free_variables bound value acc) acc fields + | Typed.TField (record, _) -> free_variables bound record acc + | Typed.TFilter (collection, predicate) -> + free_variables bound predicate (free_variables bound collection acc) + | Typed.TMap (collection, projection) -> + free_variables bound projection (free_variables bound collection acc) + | Typed.TSum collection | Typed.TCount collection -> free_variables bound collection acc + +let program typed = + let ctx = + make_context + (List.map (fun helper -> (Ident.stamp helper.Typed.th_ident, helper.Typed.th_body)) typed.Typed.tp_helpers) + 20000 + in + let helpers = + List.map + (fun helper -> { helper with Typed.th_body = specialize ctx helper.Typed.th_body }) + typed.Typed.tp_helpers + in + let query_body = specialize ctx typed.Typed.tp_query_body in + (match ctx.error with + | Some (span, message) -> Diagnostic.error span "%s" message + | None -> ()); + { typed with Typed.tp_helpers = helpers; tp_query_body = query_body } diff --git a/src/specialize.mli b/src/specialize.mli new file mode 100644 index 0000000..d863cb3 --- /dev/null +++ b/src/specialize.mli @@ -0,0 +1 @@ +val program : Typed.program -> Typed.program diff --git a/test/test_main.ml b/test/test_main.ml index 96ee292..68252b6 100644 --- a/test/test_main.ml +++ b/test/test_main.ml @@ -374,6 +374,7 @@ let () = Test_harness.run_suite "program" program_cases; Test_harness.run_suite "resolve" resolve_cases; Test_harness.run_suite "types" Test_type.cases; + Test_harness.run_suite "specialize" Test_type.specialize_cases; Test_harness.run_suite "interpret" Test_incremental.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) diff --git a/test/test_type.ml b/test/test_type.ml index c703945..91177cd 100644 --- a/test/test_type.ml +++ b/test/test_type.ml @@ -209,3 +209,118 @@ let cases = check_equal_string "query type" "collection int" (Types.pp program.Typed.tp_query_body.Typed.ty) ); ] + +let rec binder_stamps expr acc = + let walk = binder_stamps in + match expr.Typed.te with + | Typed.TLet (ident, bound, body) -> walk body (walk bound (Ident.stamp ident :: acc)) + | Typed.TLambda (ident, body) -> walk body (Ident.stamp ident :: acc) + | Typed.TApp (fn, argument) -> walk argument (walk fn acc) + | Typed.TIf (a, b, c) -> walk c (walk b (walk a acc)) + | Typed.TBinop (_, a, b) -> walk b (walk a acc) + | Typed.TTuple items -> List.fold_left (fun acc item -> walk item acc) acc items + | Typed.TRecord (_, fields) -> List.fold_left (fun acc (_, value) -> walk value acc) acc fields + | Typed.TField (record, _) -> walk record acc + | Typed.TFilter (collection, predicate) -> walk predicate (walk collection acc) + | Typed.TMap (collection, projection) -> walk projection (walk collection acc) + | Typed.TSum collection | Typed.TCount collection -> walk collection acc + | Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TVar _ | Typed.TSource _ -> acc + +let stamp_unique stamps = + let seen = Hashtbl.create 32 in + let rec check = function + | [] -> true + | stamp :: rest -> if Hashtbl.mem seen stamp then false else (Hashtbl.add seen stamp (); check rest) + in + check stamps + +let rec collect_map_functions expr acc = + match expr.Typed.te with + | Typed.TMap (collection, projection) -> + collect_map_functions collection (("map", projection) :: acc) + | Typed.TFilter (collection, predicate) -> + collect_map_functions collection (("filter", predicate) :: acc) + | Typed.TSum collection | Typed.TCount collection -> collect_map_functions collection acc + | Typed.TLet (_, bound, body) -> collect_map_functions body (collect_map_functions bound acc) + | _ -> acc + +let function_is_lambda expr = + match expr.Typed.te with Typed.TLambda _ -> true | _ -> false + +let rec has_let expr = + match expr.Typed.te with + | Typed.TLet (_, _, _) -> true + | Typed.TLambda (_, body) -> has_let body + | Typed.TApp (fn, argument) -> has_let fn || has_let argument + | Typed.TIf (a, b, c) -> has_let a || has_let b || has_let c + | Typed.TBinop (_, a, b) -> has_let a || has_let b + | Typed.TTuple items -> List.exists has_let items + | Typed.TRecord (_, fields) -> List.exists (fun (_, value) -> has_let value) fields + | Typed.TField (record, _) -> has_let record + | Typed.TFilter (collection, predicate) -> has_let collection || has_let predicate + | Typed.TMap (collection, projection) -> has_let collection || has_let projection + | Typed.TSum collection | Typed.TCount collection -> has_let collection + | Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TVar _ | Typed.TSource _ -> false + +let specialize text = Specialize.program (infer text) + +let specialize_cases = + [ + ( "inlining a helper twice does not duplicate binder stamps", + fun () -> + Ident.reset (); + let program = + specialize + "input rows : collection int\nlet scale n = n * 2\nlet quad n = scale (scale n)\nquery q = rows |> map (fun r -> quad r + quad r)\n" + in + check "all binders are distinct" + (stamp_unique (binder_stamps program.Typed.tp_query_body [])) ); + ( "a helper used as a mapping function specializes to a lambda", + fun () -> + let program = + specialize + "type order = { total : int }\ninput rows : collection order\nlet value o = o.total\nquery q = rows |> map value\n" + in + let functions = collect_map_functions program.Typed.tp_query_body [] in + check "map function is a lambda" + (List.for_all (fun (_, fn) -> function_is_lambda fn) functions) ); + ( "a helper used as a predicate specializes to a lambda", + fun () -> + let program = + specialize + "type order = { total : int }\ninput rows : collection order\nlet big o = o.total > 1000\nquery q = rows |> filter big |> count\n" + in + let functions = collect_map_functions program.Typed.tp_query_body [] in + check_equal_int "one collection function" 1 (List.length functions); + check "filter predicate is a lambda" + (List.for_all (fun (_, fn) -> function_is_lambda fn) functions) ); + ( "higher order helpers specialize at their call sites", + fun () -> + let program = + specialize + "input rows : collection int\nlet twice f x = f (f x)\nlet inc n = n + 1\nquery q = rows |> map (fun r -> twice inc r) |> sum\n" + in + check_equal_int "no helper references remain" 0 + (List.length (collect_map_functions program.Typed.tp_query_body []) - 1) ); + ( "an escaping function value is rejected", + fun () -> + check_message "escaping function" + "the mapping function must be a function literal or a helper that specializes to one, but this expression is a function value" + (fun () -> + specialize + "input rows : collection int\nlet choose b = if b then (fun n -> n) else (fun n -> n + 1)\nquery q = rows |> map (choose true)\n") ); + ( "a predicate that uses the input collection is rejected", + fun () -> + check_message "captured collection" + "the predicate cannot be incrementalized because it uses `rows` from the surrounding scope" + (fun () -> + specialize + "input rows : collection int\nquery q = rows |> filter (fun r -> r > count rows) |> count\n") ); + ( "plain let bindings are inlined into the query", + fun () -> + let program = + specialize + "input rows : collection int\nquery q = let base = 10 in rows |> map (fun r -> r + base) |> sum\n" + in + check "no let bindings survive" (not (has_let program.Typed.tp_query_body)) ); + ]