diff --git a/examples/dl-activations/relu3.py b/examples/dl-activations/relu3.py index 992817283..7dd9fb5e9 100644 --- a/examples/dl-activations/relu3.py +++ b/examples/dl-activations/relu3.py @@ -63,6 +63,18 @@ def vrelu3_embedded_ks_checkpointed_map(): ) +embedded_cpp_entry_points = """ +namespace ks { +ks::tensor<1, ks::Float> entry(ks::allocator * $alloc, ks::tensor<1, ks::Float> t) { + return ks::vrelu3($alloc, t); +} +ks::tensor<1, ks::Float> entry_vjp(ks::allocator * $alloc, ks::tensor<1, ks::Float> t, ks::tensor<1, ks::Float> dret) { + return ks::sufrev_vrelu3($alloc, t, dret); +} +} +""" + + def vrelu3_embedded_cpp_inlined_map(): return cpp_string_to_autograd_function( """ @@ -113,10 +125,9 @@ def vrelu3_embedded_cpp_inlined_map(): return ret; } } - """, - "vrelu3", + """ + + embedded_cpp_entry_points, "ksc_dl_activations__manual__vrelu3_embedded_cpp_inlined_map", - generate_lm=False, ) @@ -161,10 +172,9 @@ def vrelu3_embedded_cpp_mask(): return ret; } } - """, - "vrelu3", + """ + + embedded_cpp_entry_points, "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask", - generate_lm=False, ) @@ -214,10 +224,9 @@ def vrelu3_embedded_cpp_mask_bool_to_float(): return ret; } } - """, - "vrelu3", + """ + + embedded_cpp_entry_points, "ksc_dl_activations__manual__vrelu3_embedded_cpp_mask_bool_to_float", - generate_lm=False, ) diff --git a/src/python/ksc/torch_frontend.py b/src/python/ksc/torch_frontend.py index c433371f6..a65e99860 100644 --- a/src/python/ksc/torch_frontend.py +++ b/src/python/ksc/torch_frontend.py @@ -458,19 +458,18 @@ def forward_template(py_mod, ctx, *args): return torch_from_ks(outputs) -def backward_template(py_mod, generate_lm, ctx, *args): +def backward_template(py_mod, ctx, *args): ks_args = make_tuple_if_many_args(torch_to_ks(py_mod, x) for x in ctx.saved_tensors) ks_grad_args = make_tuple_if_many_args(torch_to_ks(py_mod, x) for x in args) - rev_entry = py_mod.rev_entry if generate_lm else py_mod.sufrev_entry - outputs = rev_entry(ks_args, ks_grad_args) + outputs = py_mod.entry_vjp(ks_args, ks_grad_args) return torch_from_ks(outputs) -def make_KscAutogradFunction(py_mod, generate_lm): +def make_KscAutogradFunction(py_mod): # We need to make a new class for every py_mod, as PyTorch requires forward and backward to be # staticmethods. This is not too expensive, as each mod needs to be compiled anyway. forward = lambda ctx, args: forward_template(py_mod, ctx, args) - backward = lambda ctx, args: backward_template(py_mod, generate_lm, ctx, args) + backward = lambda ctx, args: backward_template(py_mod, ctx, args) return type( "KscAutogradFunction_" + py_mod.__name__, (torch.autograd.Function,), @@ -483,9 +482,7 @@ def make_KscAutogradFunction(py_mod, generate_lm): ) -def ksc_defs_to_module( - ksc_defs, entry_def, derivatives_to_generate, torch_extension_name -): +def ksc_defs_to_module(ksc_defs, entry_def, torch_extension_name, generate_lm): symtab = dict() ksc_dir = utils.get_ksc_dir() decls_prelude = list(parse_ks_filename(ksc_dir + "/src/runtime/prelude.ks")) @@ -503,48 +500,40 @@ def ksc_defs_to_module( defs_with_derivatives = [] for ksc_def in ksc_defs: defs_with_derivatives += [ksc_def] - if "sufrev" in derivatives_to_generate: + if generate_lm: + defs_with_derivatives += [ + GDef("rev", ksc_def.name), + ] + else: defs_with_derivatives += [ GDef("suffwdpass", ksc_def.name), GDef("sufrevpass", ksc_def.name), GDef("sufrev", ksc_def.name), ] - if "fwd" in derivatives_to_generate: - defs_with_derivatives += [ - GDef("fwd", ksc_def.name), - ] - if "rev" in derivatives_to_generate: - defs_with_derivatives += [ - GDef("rev", ksc_def.name), - ] ks_str = "\n".join(map(pformat, defs_with_derivatives)) return ksc_string_to_module( - ks_str, entry_def.name, derivatives_to_generate, torch_extension_name + ks_str, entry_def.name, torch_extension_name, generate_lm ) -def ksc_string_to_module( - ks_str, entry_sn, derivatives_to_generate, torch_extension_name -): - bindings_to_generate = [("entry", entry_sn)] + [ - (f"{der}_entry", StructuredName((der, entry_sn))) - for der in derivatives_to_generate +def ksc_string_to_module(ks_str, entry_sn, torch_extension_name, generate_lm): + 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 ) -def cpp_string_to_module( - cpp_str, entry_name, derivatives_to_generate, torch_extension_name -): - bindings_to_generate = [("entry", entry_name)] + [ - (f"{der}_entry", f"{der}_{entry_name}") for der in derivatives_to_generate +def cpp_string_to_module(cpp_str, torch_extension_name, entry_name, entry_vjp_name): + 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, ) @@ -553,31 +542,24 @@ def cpp_string_to_module( def ksc_defs_to_autograd_function( ksc_defs, entry_def, torch_extension_name, generate_lm=True ): - derivatives_to_generate = ["fwd", "rev"] if generate_lm else ["sufrev"] - mod = ksc_defs_to_module( - ksc_defs, entry_def, derivatives_to_generate, torch_extension_name - ) - return make_KscAutogradFunction(mod, generate_lm) + mod = ksc_defs_to_module(ksc_defs, entry_def, torch_extension_name, generate_lm) + return make_KscAutogradFunction(mod) def ksc_string_to_autograd_function( - ks_str, entry_sn, torch_extension_name, generate_lm + ks_str, entry_sn, torch_extension_name, generate_lm=True ): - derivatives_to_generate = ["fwd", "rev"] if generate_lm else ["sufrev"] - mod = ksc_string_to_module( - ks_str, entry_sn, derivatives_to_generate, torch_extension_name - ) - return make_KscAutogradFunction(mod, generate_lm) + mod = ksc_string_to_module(ks_str, entry_sn, torch_extension_name, generate_lm) + return make_KscAutogradFunction(mod) def cpp_string_to_autograd_function( - cpp_str, entry_name, torch_extension_name, generate_lm + cpp_str, torch_extension_name, entry_name="entry", entry_vjp_name="entry_vjp", ): - derivatives_to_generate = ["fwd", "rev"] if generate_lm else ["sufrev"] mod = cpp_string_to_module( - cpp_str, entry_name, derivatives_to_generate, torch_extension_name + cpp_str, torch_extension_name, entry_name, entry_vjp_name ) - return make_KscAutogradFunction(mod, generate_lm) + return make_KscAutogradFunction(mod) import inspect diff --git a/test/ts2k/test_ts2k.py b/test/ts2k/test_ts2k.py index 67e977785..aa2d05c73 100644 --- a/test/ts2k/test_ts2k.py +++ b/test/ts2k/test_ts2k.py @@ -73,7 +73,7 @@ def test_ts2k_relux(): def test_ts2k_relux_grad(): compile_relux() - ks_ans = ks_relux.py_mod.rev_entry(1.3, 1.0) + ks_ans = ks_relux.py_mod.entry_vjp(1.3, 1.0) ans = grad_relux(1.3) assert pytest.approx(ks_ans, 1e-6) == ans @@ -116,7 +116,7 @@ def test_bar(): assert pytest.approx(ks_ans, 1e-5) == ans # Check grad - ks_ans = ks_bar.py_mod.rev_entry((a, x), 1.0) + ks_ans = ks_bar.py_mod.entry_vjp((a, x), 1.0) ans = grad_bar(a, x) assert pytest.approx(ks_ans[1], 1e-5) == ans[1] @@ -195,10 +195,7 @@ def test_relu3(generate_lm): # Test gradient ks == py py_ans = grad_relu3(x) - grad_fun = ( - ks_relu3.py_mod.rev_entry if generate_lm else ks_relu3.py_mod.sufrev_entry - ) - ks_ans = grad_fun(x, 1.0) + ks_ans = ks_relu3.py_mod.entry_vjp(x, 1.0) assert pytest.approx(ks_ans, 1e-6) == py_ans