From 4b1265d337d712522d2ee1bf964e98dd5759b2aa Mon Sep 17 00:00:00 2001 From: Tom Ellis Date: Mon, 12 Jul 2021 13:04:58 +0100 Subject: [PATCH 1/3] WIP allow Cflags --- examples/dl-activations/relu3.py | 17 ++++++++++++ src/python/ksc/compile.py | 28 +++++--------------- src/python/ksc/torch_frontend.py | 45 +++++++++++++++++++++++++------- 3 files changed, 59 insertions(+), 31 deletions(-) diff --git a/examples/dl-activations/relu3.py b/examples/dl-activations/relu3.py index bf63b4d8c..8d0f3c8d5 100644 --- a/examples/dl-activations/relu3.py +++ b/examples/dl-activations/relu3.py @@ -55,6 +55,14 @@ def vrelu3(x: torch.Tensor): return elementwise_apply_hack("relu3", x) +embedded_cflags = [ + "-std=c++17", + "-g", + "-O3", + # "-DKS_BOUNDS_CHECK", +] + + def vrelu3_embedded_ks_checkpointed_map(): return ksc_string_to_autograd_function( """(def relu3 Float (x : Float) @@ -80,6 +88,7 @@ def vrelu3_embedded_ks_checkpointed_map(): expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))), "ksc_dl_activations__manual__vrelu3_embedded_ks_checkpointed_map", generate_lm=False, + extra_cflags=embedded_cflags, ) @@ -156,6 +165,7 @@ def vrelu3_embedded_cpp_inlined_map(): """ + embedded_cpp_entry_points, "ksc_dl_activations__manual__vrelu3_embedded_cpp_inlined_map", + extra_cflags=embedded_cflags, ) @@ -203,6 +213,7 @@ def vrelu3_embedded_cpp_mask(): """ + embedded_cpp_entry_points, "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask", + extra_cflags=embedded_cflags, ) @@ -255,6 +266,7 @@ def vrelu3_embedded_cpp_mask_bool_to_float(): """ + embedded_cpp_entry_points, "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask_bool_to_float", + extra_cflags=embedded_cflags, ) @@ -286,6 +298,7 @@ def vrelu3_embedded_ks_checkpointed_map_handwritten_relu3(): expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))), "ksc_dl_activations__manual__vrelu3_embedded_ks_checkpointed_map_handwritten_relu3", generate_lm=False, + extra_cflags=embedded_cflags, ) @@ -313,6 +326,7 @@ def vrelu3_embedded_ks_checkpointed_map_handwritten_inlined_relu3(): expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))), "ksc_dl_activations__manual__vrelu3_embedded_ks_checkpointed_map_handwritten_inlined_relu3", generate_lm=False, + extra_cflags=embedded_cflags, ) @@ -343,6 +357,7 @@ def vrelu3_embedded_ks_checkpointed_map_mask(): expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))), "ksc_dl_activations__manual__vrelu3_embedded_ks_checkpointed_map_mask", generate_lm=False, + extra_cflags=embedded_cflags, ) @@ -364,6 +379,7 @@ def vrelu3_embedded_INCORRECT_ks_upper_bound_via_map(): expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))), "ksc_dl_activations__manual__vrelu3_embedded_INCORRECT_ks_upper_bound_via_map", generate_lm=False, + extra_cflags=embedded_cflags, ) @@ -381,6 +397,7 @@ def vrelu3_embedded_INCORRECT_ks_upper_bound(): expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))), "ksc_dl_activations__manual__vrelu3_embedded_INCORRECT_ks_upper_bound", generate_lm=False, + extra_cflags=embedded_cflags, ) diff --git a/src/python/ksc/compile.py b/src/python/ksc/compile.py index 6d449de22..04dfa053f 100644 --- a/src/python/ksc/compile.py +++ b/src/python/ksc/compile.py @@ -224,7 +224,7 @@ def build_py_module_from_ks( def build_module_using_pytorch_from_ks( - ks_str, bindings_to_generate, torch_extension_name, use_aten=False + ks_str, bindings_to_generate, torch_extension_name, use_aten=False, extra_cflags=[] ): """Uses PyTorch C++ extension mechanism to build and load a module @@ -247,43 +247,27 @@ def build_module_using_pytorch_from_ks( ) return build_module_using_pytorch_from_cpp_backend( - cpp_str, torch_extension_name, use_aten + cpp_str, torch_extension_name, use_aten, extra_cflags ) def build_module_using_pytorch_from_cpp( - cpp_str, bindings_to_generate, torch_extension_name, use_aten + cpp_str, bindings_to_generate, torch_extension_name, use_aten, extra_cflags=[] ): cpp_pybind = generate_cpp_pybind_module_declaration( bindings_to_generate, torch_extension_name ) return build_module_using_pytorch_from_cpp_backend( - cpp_str + cpp_pybind, torch_extension_name, use_aten + cpp_str + cpp_pybind, torch_extension_name, use_aten, extra_cflags ) def build_module_using_pytorch_from_cpp_backend( - cpp_str, torch_extension_name, use_aten + cpp_str, torch_extension_name, use_aten, extra_cflags ): __ksc_path, ksc_runtime_dir = utils.get_ksc_paths() - extra_cflags = ["-DKS_INCLUDE_ATEN"] if use_aten else [] - - # I don't like this assumption about Windows -> cl but it matches what PyTorch is currently doing: - # https://github.com/pytorch/pytorch/blob/ad8d1b2aaaf2ba28c51b1cb38f86311749eff755/torch/utils/cpp_extension.py#L1374-L1378 - # We're making a guess here if people recognifigure their C++ compiler on Windows it's because they're using non-MSVC - # otherwise we need to inspect the end of the path path for cl[.exe]. - - cpp_compiler = os.environ.get("CXX") - if cpp_compiler == None and sys.platform == "win32": - extra_cflags += ["/std:c++17", "/O2"] - else: - extra_cflags += [ - "-std=c++17", - "-g", - "-O3", - # "-DKS_BOUNDS_CHECK", - ] + extra_cflags += ["-DKS_INCLUDE_ATEN"] if use_aten else [] verbose = True diff --git a/src/python/ksc/torch_frontend.py b/src/python/ksc/torch_frontend.py index c022dd2fd..c3326bded 100644 --- a/src/python/ksc/torch_frontend.py +++ b/src/python/ksc/torch_frontend.py @@ -498,28 +498,49 @@ def ksc_defs_to_module(ksc_defs, entry_def, torch_extension_name, generate_lm): ks_str = "\n".join(map(pformat, defs_with_derivatives)) return ksc_string_to_module( - ks_str, entry_def.name, torch_extension_name, generate_lm + ks_str, + entry_def.name, + torch_extension_name, + generate_lm, + extra_cflags=[ + "-std=c++17", + "-g", + "-O3", + # "-DKS_BOUNDS_CHECK", + ], ) -def ksc_string_to_module(ks_str, entry_sn, torch_extension_name, generate_lm): +def ksc_string_to_module( + ks_str, entry_sn, torch_extension_name, generate_lm, extra_cflags +): der = "rev" if generate_lm else "sufrev" bindings_to_generate = [ ("entry", entry_sn), ("entry_vjp", StructuredName((der, entry_sn))), ] return build_module_using_pytorch_from_ks( - ks_str, bindings_to_generate, torch_extension_name, use_aten=True + ks_str, + bindings_to_generate, + torch_extension_name, + use_aten=True, + extra_cflags=extra_cflags, ) -def cpp_string_to_module(cpp_str, torch_extension_name, entry_name, entry_vjp_name): +def cpp_string_to_module( + cpp_str, torch_extension_name, entry_name, entry_vjp_name, extra_cflags +): bindings_to_generate = [ ("entry", entry_name), ("entry_vjp", entry_vjp_name), ] return build_module_using_pytorch_from_cpp( - cpp_str, bindings_to_generate, torch_extension_name, use_aten=True, + cpp_str, + bindings_to_generate, + torch_extension_name, + use_aten=True, + extra_cflags=extra_cflags, ) @@ -531,17 +552,23 @@ def ksc_defs_to_autograd_function( def ksc_string_to_autograd_function( - ks_str, entry_sn, torch_extension_name, generate_lm=True + ks_str, entry_sn, torch_extension_name, generate_lm=True, extra_cflags=[] ): - mod = ksc_string_to_module(ks_str, entry_sn, torch_extension_name, generate_lm) + mod = ksc_string_to_module( + ks_str, entry_sn, torch_extension_name, generate_lm, extra_cflags + ) return make_KscAutogradFunction(mod) def cpp_string_to_autograd_function( - cpp_str, torch_extension_name, entry_name="entry", entry_vjp_name="entry_vjp", + cpp_str, + torch_extension_name, + entry_name="entry", + entry_vjp_name="entry_vjp", + extra_cflags=[], ): mod = cpp_string_to_module( - cpp_str, torch_extension_name, entry_name, entry_vjp_name + cpp_str, torch_extension_name, entry_name, entry_vjp_name, extra_cflags ) return make_KscAutogradFunction(mod) From b59dd70fb1ee66dd1938667fd5942ec968418b02 Mon Sep 17 00:00:00 2001 From: Tom Ellis Date: Mon, 12 Jul 2021 13:23:53 +0100 Subject: [PATCH 2/3] Add new --- examples/dl-activations/relu3.py | 294 ++++++++++++++++++++----------- 1 file changed, 188 insertions(+), 106 deletions(-) diff --git a/examples/dl-activations/relu3.py b/examples/dl-activations/relu3.py index 8d0f3c8d5..84ae35b6d 100644 --- a/examples/dl-activations/relu3.py +++ b/examples/dl-activations/relu3.py @@ -63,58 +63,15 @@ def vrelu3(x: torch.Tensor): ] -def vrelu3_embedded_ks_checkpointed_map(): - return ksc_string_to_autograd_function( - """(def relu3 Float (x : Float) - (if (lt x 0.0) - 0.0 - (if (lt x 1.0) - (div (mul x (mul x x)) 3.0) - (sub x (div 2.0 3.0))))) - - (gdef suffwdpass [relu3 Float]) - (gdef sufrevpass [relu3 Float]) - (gdef sufrev [relu3 Float]) - - (def [vrelu3 (Vec Float)] (Vec Float) - (t : Vec Float) - (map (lam (ti : Float) (relu3 ti)) t)) - - (def [sufrev [vrelu3 (Vec Float)]] (Vec Float) - ((t : Vec Float) (dret : Vec Float)) - ; TODO: 1.0 should be dret[i] - luckily we are called with dret==1.0 - (map (lam (ti : Float) ([sufrev [relu3 Float]] ti 1.0)) t)) - """, - expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))), - "ksc_dl_activations__manual__vrelu3_embedded_ks_checkpointed_map", - generate_lm=False, - extra_cflags=embedded_cflags, - ) - - -embedded_cpp_entry_points = """ -#include "knossos-entry-points-torch.h" - -torch::Tensor entry(torch::Tensor t) { - using namespace ks::entry_points; - auto ks_t = convert_argument>(t); - auto ks_ret = ks::vrelu3(&g_alloc, ks_t); - return convert_return_value(ks_ret); -} - -torch::Tensor entry_vjp(torch::Tensor t, torch::Tensor dret) { - using namespace ks::entry_points; - auto ks_t = convert_argument>(t); - auto ks_dret = convert_argument>(dret); - auto ks_ret = ks::sufrev_vrelu3(&g_alloc, ks_t, ks_dret); - return convert_return_value(ks_ret); -} -""" +embedded_cflags_opts = [ + "-march=native", + "-funroll-loops", + "-ffast-math", + "-mprefer-vector-width=512", +] -def vrelu3_embedded_cpp_inlined_map(): - return cpp_string_to_autograd_function( - """ +cpp_mask_bool_to_float = """ #include "knossos.h" namespace ks{ @@ -122,19 +79,17 @@ def vrelu3_embedded_cpp_inlined_map(): auto tdata = t.data(); auto ret = tensor<1, ks::Float>::create($alloc, t.size()); auto retdata = ret.data(); + ks::Float third = 1.0 / 3.0; + ks::Float two_thirds = 2.0 / 3.0; for (int i = 0, ne = t.num_elements(); i != ne; ++i) { - ks::Float c$1; ks::Float x = tdata[i]; - if (x < 0.0) { - c$1 = 0.0; - } else { - if (x < 1.0) { - c$1 = x * x * x / 3.0; - } else { - c$1 = x - 2.0 / 3.0; - } - } - retdata[i] = c$1; + auto val0to1 = third * x * x * x; + auto val1up = x - two_thirds; + auto le1 = x <= 1; + + retdata[i] = bool_to_float$ab($alloc, x > 0) + * (bool_to_float$ab($alloc, le1) * val0to1 + + bool_to_float$ab($alloc, !le1) * val1up); } return ret; } @@ -145,33 +100,24 @@ def vrelu3_embedded_cpp_inlined_map(): auto ret = tensor<1, ks::Float>::create($alloc, t.size()); auto retdata = ret.data(); for (int i = 0, ne = t.num_elements(); i != ne; ++i) { - ks::Float c$1; ks::Float x = tdata[i]; ks::Float dreti = dretdata[i]; - if (x < 0.0) { - c$1 = 0.0; - } else { - if (x < 1.0) { - c$1 = x * x; - } else { - c$1 = 1.0; - } - } - retdata[i] = c$1 * dreti; + auto val0to1 = x * x; + auto val1up = 1.0; + + auto le1 = x <= 1; + + retdata[i] = bool_to_float$ab($alloc, x > 0) + * (bool_to_float$ab($alloc, le1) * val0to1 + + bool_to_float$ab($alloc, !le1) * val1up) + * dreti; } return ret; } } """ - + embedded_cpp_entry_points, - "ksc_dl_activations__manual__vrelu3_embedded_cpp_inlined_map", - extra_cflags=embedded_cflags, - ) - -def vrelu3_embedded_cpp_mask(): - return cpp_string_to_autograd_function( - """ +cpp_mask = """ #include "knossos.h" namespace ks{ @@ -211,15 +157,9 @@ def vrelu3_embedded_cpp_mask(): } } """ - + embedded_cpp_entry_points, - "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask", - extra_cflags=embedded_cflags, - ) -def vrelu3_embedded_cpp_mask_bool_to_float(): - return cpp_string_to_autograd_function( - """ +cpp_inlined_map = """ #include "knossos.h" namespace ks{ @@ -227,17 +167,19 @@ def vrelu3_embedded_cpp_mask_bool_to_float(): auto tdata = t.data(); auto ret = tensor<1, ks::Float>::create($alloc, t.size()); auto retdata = ret.data(); - ks::Float third = 1.0 / 3.0; - ks::Float two_thirds = 2.0 / 3.0; for (int i = 0, ne = t.num_elements(); i != ne; ++i) { + ks::Float c$1; ks::Float x = tdata[i]; - auto val0to1 = third * x * x * x; - auto val1up = x - two_thirds; - auto le1 = x <= 1; - - retdata[i] = bool_to_float$ab($alloc, x > 0) - * (bool_to_float$ab($alloc, le1) * val0to1 - + bool_to_float$ab($alloc, !le1) * val1up); + if (x < 0.0) { + c$1 = 0.0; + } else { + if (x < 1.0) { + c$1 = x * x * x / 3.0; + } else { + c$1 = x - 2.0 / 3.0; + } + } + retdata[i] = c$1; } return ret; } @@ -248,28 +190,168 @@ def vrelu3_embedded_cpp_mask_bool_to_float(): auto ret = tensor<1, ks::Float>::create($alloc, t.size()); auto retdata = ret.data(); for (int i = 0, ne = t.num_elements(); i != ne; ++i) { + ks::Float c$1; ks::Float x = tdata[i]; ks::Float dreti = dretdata[i]; - auto val0to1 = x * x; - auto val1up = 1.0; - - auto le1 = x <= 1; - - retdata[i] = bool_to_float$ab($alloc, x > 0) - * (bool_to_float$ab($alloc, le1) * val0to1 - + bool_to_float$ab($alloc, !le1) * val1up) - * dreti; + if (x < 0.0) { + c$1 = 0.0; + } else { + if (x < 1.0) { + c$1 = x * x; + } else { + c$1 = 1.0; + } + } + retdata[i] = c$1 * dreti; } return ret; } } """ - + embedded_cpp_entry_points, + + +def vrelu3_embedded_ks_checkpointed_map(): + return ksc_string_to_autograd_function( + """(def relu3 Float (x : Float) + (if (lt x 0.0) + 0.0 + (if (lt x 1.0) + (div (mul x (mul x x)) 3.0) + (sub x (div 2.0 3.0))))) + + (gdef suffwdpass [relu3 Float]) + (gdef sufrevpass [relu3 Float]) + (gdef sufrev [relu3 Float]) + + (def [vrelu3 (Vec Float)] (Vec Float) + (t : Vec Float) + (map (lam (ti : Float) (relu3 ti)) t)) + + (def [sufrev [vrelu3 (Vec Float)]] (Vec Float) + ((t : Vec Float) (dret : Vec Float)) + ; TODO: 1.0 should be dret[i] - luckily we are called with dret==1.0 + (map (lam (ti : Float) ([sufrev [relu3 Float]] ti 1.0)) t)) + """, + expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))), + "ksc_dl_activations__manual__vrelu3_embedded_ks_checkpointed_map", + generate_lm=False, + extra_cflags=embedded_cflags, + ) + + +def vrelu3_embedded_ks_checkpointed_map_flags(): + return ksc_string_to_autograd_function( + """(def relu3 Float (x : Float) + (if (lt x 0.0) + 0.0 + (if (lt x 1.0) + (div (mul x (mul x x)) 3.0) + (sub x (div 2.0 3.0))))) + + (gdef suffwdpass [relu3 Float]) + (gdef sufrevpass [relu3 Float]) + (gdef sufrev [relu3 Float]) + + (def [vrelu3 (Vec Float)] (Vec Float) + (t : Vec Float) + (map (lam (ti : Float) (relu3 ti)) t)) + + (def [sufrev [vrelu3 (Vec Float)]] (Vec Float) + ((t : Vec Float) (dret : Vec Float)) + ; TODO: 1.0 should be dret[i] - luckily we are called with dret==1.0 + (map (lam (ti : Float) ([sufrev [relu3 Float]] ti 1.0)) t)) + """, + expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))), + "ksc_dl_activations__manual__vrelu3_embedded_ks_checkpointed_map_flags", + generate_lm=False, + extra_cflags=embedded_cflags + embedded_cflags_opts, + ) + + +embedded_cpp_entry_points = """ +#include "knossos-entry-points-torch.h" + +torch::Tensor entry(torch::Tensor t) { + using namespace ks::entry_points; + auto ks_t = convert_argument>(t); + auto ks_ret = ks::vrelu3(&g_alloc, ks_t); + return convert_return_value(ks_ret); +} + +torch::Tensor entry_vjp(torch::Tensor t, torch::Tensor dret) { + using namespace ks::entry_points; + auto ks_t = convert_argument>(t); + auto ks_dret = convert_argument>(dret); + auto ks_ret = ks::sufrev_vrelu3(&g_alloc, ks_t, ks_dret); + return convert_return_value(ks_ret); +} +""" + + +def vrelu3_embedded_cpp_inlined_map(): + return cpp_string_to_autograd_function( + cpp_inlined_map + embedded_cpp_entry_points, + "ksc_dl_activations__manual__vrelu3_embedded_cpp_inlined_map", + extra_cflags=embedded_cflags, + ) + + +def vrelu3_embedded_cpp_inlined_map_flags(): + return cpp_string_to_autograd_function( + cpp_inlined_map + embedded_cpp_entry_points, + "ksc_dl_activations__manual__vrelu3_embedded_cpp_inlined_map_flags", + extra_cflags=embedded_cflags + embedded_cflags_opts, + ) + + +def vrelu3_embedded_cpp_mask(): + return cpp_string_to_autograd_function( + cpp_mask + embedded_cpp_entry_points, + "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask", + extra_cflags=embedded_cflags, + ) + + +def vrelu3_embedded_cpp_mask_bool_to_float(): + return cpp_string_to_autograd_function( + cpp_mask_bool_to_float + embedded_cpp_entry_points, "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask_bool_to_float", extra_cflags=embedded_cflags, ) +def vrelu3_embedded_cpp_mask_flags(): + return cpp_string_to_autograd_function( + cpp_mask + embedded_cpp_entry_points, + "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask_flags", + extra_cflags=embedded_cflags + embedded_cflags_opts, + ) + + +def vrelu3_embedded_cpp_mask_flags_tune(): + return cpp_string_to_autograd_function( + cpp_mask + embedded_cpp_entry_points, + "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask_flags_tune", + extra_cflags=embedded_cflags + embedded_cflags_opts + ["-mtune=native"], + ) + + +def vrelu3_embedded_cpp_mask_bool_to_float_flags(): + return cpp_string_to_autograd_function( + cpp_mask_bool_to_float + embedded_cpp_entry_points, + "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask_bool_to_float_flags", + extra_cflags=embedded_cflags + embedded_cflags_opts, + ) + + +def vrelu3_embedded_cpp_mask_bool_to_float_flags_tune(): + return cpp_string_to_autograd_function( + cpp_mask_bool_to_float + embedded_cpp_entry_points, + "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask_bool_to_float_flags_tune", + extra_cflags=embedded_cflags + embedded_cflags_opts + ["-mtune=native"], + ) + + def vrelu3_embedded_ks_checkpointed_map_handwritten_relu3(): return ksc_string_to_autograd_function( """(def relu3 Float (x : Float) From 1e9325214d992d3af61d8bc2a6e20e8174ee4df1 Mon Sep 17 00:00:00 2001 From: Tom Ellis Date: Thu, 15 Jul 2021 15:38:20 +0100 Subject: [PATCH 3/3] Add CFlags dataclass --- examples/dl-activations/relu3.py | 23 ++++++--------- src/python/ksc/compile.py | 50 +++++++++++++++++++++++++++++++- src/python/ksc/torch_frontend.py | 16 +++++----- 3 files changed, 66 insertions(+), 23 deletions(-) diff --git a/examples/dl-activations/relu3.py b/examples/dl-activations/relu3.py index 84ae35b6d..0b173137b 100644 --- a/examples/dl-activations/relu3.py +++ b/examples/dl-activations/relu3.py @@ -3,6 +3,7 @@ import os from collections import OrderedDict from ksc import utils +import ksc.compile import ksc.expr as expr from ksc.type import Type from ksc.torch_frontend import ( @@ -55,21 +56,15 @@ def vrelu3(x: torch.Tensor): return elementwise_apply_hack("relu3", x) -embedded_cflags = [ - "-std=c++17", - "-g", - "-O3", - # "-DKS_BOUNDS_CHECK", -] +embedded_cflags = ksc.compile.default_cflags -embedded_cflags_opts = [ - "-march=native", - "-funroll-loops", - "-ffast-math", - "-mprefer-vector-width=512", -] +embedded_cflags_opts = ksc.compile.CFlags.GCCOnly( + ["-march=native", "-funroll-loops", "-ffast-math", "-mprefer-vector-width=512",] +) + +mtune_cflags = ksc.compile.CFlags.GCCOnly(["-mtune=native"]) cpp_mask_bool_to_float = """ #include "knossos.h" @@ -332,7 +327,7 @@ def vrelu3_embedded_cpp_mask_flags_tune(): return cpp_string_to_autograd_function( cpp_mask + embedded_cpp_entry_points, "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask_flags_tune", - extra_cflags=embedded_cflags + embedded_cflags_opts + ["-mtune=native"], + extra_cflags=embedded_cflags + embedded_cflags_opts + mtune_cflags, ) @@ -348,7 +343,7 @@ def vrelu3_embedded_cpp_mask_bool_to_float_flags_tune(): return cpp_string_to_autograd_function( cpp_mask_bool_to_float + embedded_cpp_entry_points, "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask_bool_to_float_flags_tune", - extra_cflags=embedded_cflags + embedded_cflags_opts + ["-mtune=native"], + extra_cflags=embedded_cflags + embedded_cflags_opts + mtune_cflags, ) diff --git a/src/python/ksc/compile.py b/src/python/ksc/compile.py index 04dfa053f..50be7cd56 100644 --- a/src/python/ksc/compile.py +++ b/src/python/ksc/compile.py @@ -3,8 +3,10 @@ import subprocess import sysconfig import sys +from dataclasses import dataclass from tempfile import NamedTemporaryFile from tempfile import gettempdir +from typing import List from torch.utils import cpp_extension @@ -14,6 +16,41 @@ preserve_temporary_files = False +@dataclass(frozen=True) +class CFlags: + cl_flags: List[str] + gcc_flags: List[str] + + def __add__(self, other): + return CFlags( + cl_flags=self.cl_flags + other.cl_flags, + gcc_flags=self.gcc_flags + other.gcc_flags, + ) + + @staticmethod + def All(cflags): + return CFlags(cl_flags=cflags, gcc_flags=cflags) + + @staticmethod + def Empty(): + return CFlags.All([]) + + @staticmethod + def GCCOnly(gcc_flags): + return CFlags(cl_flags=[], gcc_flags=gcc_flags) + + +default_cflags = CFlags( + cl_flags=["/std:c++17", "/O2"], + gcc_flags=[ + "-std=c++17", + "-g", + "-O3", + # "-DKS_BOUNDS_CHECK", + ], +) + + def subprocess_run(cmd, env=None): return ( subprocess.run(cmd, stdout=subprocess.PIPE, env=env).stdout.decode().strip("\n") @@ -267,7 +304,18 @@ def build_module_using_pytorch_from_cpp_backend( ): __ksc_path, ksc_runtime_dir = utils.get_ksc_paths() - extra_cflags += ["-DKS_INCLUDE_ATEN"] if use_aten else [] + extra_cflags = extra_cflags + CFlags.All(["-DKS_INCLUDE_ATEN"] if use_aten else []) + + # I don't like this assumption about Windows -> cl but it matches what PyTorch is currently doing: + # https://github.com/pytorch/pytorch/blob/ad8d1b2aaaf2ba28c51b1cb38f86311749eff755/torch/utils/cpp_extension.py#L1374-L1378 + # We're making a guess here if people recognifigure their C++ compiler on Windows it's because they're using non-MSVC + # otherwise we need to inspect the end of the path path for cl[.exe]. + + cpp_compiler = os.environ.get("CXX") + if cpp_compiler == None and sys.platform == "win32": + extra_cflags = extra_cflags.cl_flags + else: + extra_cflags = extra_cflags.gcc_flags verbose = True diff --git a/src/python/ksc/torch_frontend.py b/src/python/ksc/torch_frontend.py index c3326bded..62b2defec 100644 --- a/src/python/ksc/torch_frontend.py +++ b/src/python/ksc/torch_frontend.py @@ -13,6 +13,7 @@ from ksc.compile import ( build_module_using_pytorch_from_ks, build_module_using_pytorch_from_cpp, + default_cflags, ) from ksc.type import Type @@ -502,12 +503,7 @@ def ksc_defs_to_module(ksc_defs, entry_def, torch_extension_name, generate_lm): entry_def.name, torch_extension_name, generate_lm, - extra_cflags=[ - "-std=c++17", - "-g", - "-O3", - # "-DKS_BOUNDS_CHECK", - ], + extra_cflags=default_cflags, ) @@ -552,7 +548,11 @@ def ksc_defs_to_autograd_function( def ksc_string_to_autograd_function( - ks_str, entry_sn, torch_extension_name, generate_lm=True, extra_cflags=[] + ks_str, + entry_sn, + torch_extension_name, + generate_lm=True, + extra_cflags=default_cflags, ): mod = ksc_string_to_module( ks_str, entry_sn, torch_extension_name, generate_lm, extra_cflags @@ -565,7 +565,7 @@ def cpp_string_to_autograd_function( torch_extension_name, entry_name="entry", entry_vjp_name="entry_vjp", - extra_cflags=[], + extra_cflags=default_cflags, ): mod = cpp_string_to_module( cpp_str, torch_extension_name, entry_name, entry_vjp_name, extra_cflags