diff --git a/examples/dl-activations/relu3.py b/examples/dl-activations/relu3.py index bf63b4d8c..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,57 +56,17 @@ def vrelu3(x: torch.Tensor): return elementwise_apply_hack("relu3", x) -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))))) +embedded_cflags = ksc.compile.default_cflags - (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, - ) - - -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); -} +embedded_cflags_opts = ksc.compile.CFlags.GCCOnly( + ["-march=native", "-funroll-loops", "-ffast-math", "-mprefer-vector-width=512",] +) -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); -} -""" +mtune_cflags = ksc.compile.CFlags.GCCOnly(["-mtune=native"]) -def vrelu3_embedded_cpp_inlined_map(): - return cpp_string_to_autograd_function( - """ +cpp_mask_bool_to_float = """ #include "knossos.h" namespace ks{ @@ -113,19 +74,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; } @@ -136,32 +95,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", - ) - -def vrelu3_embedded_cpp_mask(): - return cpp_string_to_autograd_function( - """ +cpp_mask = """ #include "knossos.h" namespace ks{ @@ -201,14 +152,9 @@ def vrelu3_embedded_cpp_mask(): } } """ - + embedded_cpp_entry_points, - "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask", - ) -def vrelu3_embedded_cpp_mask_bool_to_float(): - return cpp_string_to_autograd_function( - """ +cpp_inlined_map = """ #include "knossos.h" namespace ks{ @@ -216,17 +162,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; } @@ -237,24 +185,165 @@ 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_cflags, + ) + + +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_cflags, ) @@ -286,6 +375,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 +403,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 +434,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 +456,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 +474,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..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") @@ -224,7 +261,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,27 +284,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 [] + 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 @@ -276,14 +313,9 @@ def build_module_using_pytorch_from_cpp_backend( cpp_compiler = os.environ.get("CXX") if cpp_compiler == None and sys.platform == "win32": - extra_cflags += ["/std:c++17", "/O2"] + extra_cflags = extra_cflags.cl_flags else: - extra_cflags += [ - "-std=c++17", - "-g", - "-O3", - # "-DKS_BOUNDS_CHECK", - ] + 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 c022dd2fd..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 @@ -498,28 +499,44 @@ 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=default_cflags, ) -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 +548,27 @@ 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=default_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=default_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)