Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 64 additions & 14 deletions scripts/generate_wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -586,14 +586,21 @@ def _generate_params(node):

return ", ".join(parts)

def _generate_arguments(node):
def _generate_arguments(
node, preconverted_tensor_arg=None, preconverted_tensor_name=None
):
args = []

for arg in node.get_arguments():
if arg.spelling == "stream":
continue

if _is_optional_tensor(arg):
if (
preconverted_tensor_arg is not None
and arg.spelling == preconverted_tensor_arg.spelling
):
args.append(f"std::move({preconverted_tensor_name})")
elif _is_optional_tensor(arg):
args.append(f"OptionalTensorFromPybind11Handle({arg.spelling})")
elif _is_vector_tensor(arg):
args.append(f"VectorTensorFromPybind11Handle({arg.spelling})")
Expand All @@ -617,19 +624,40 @@ def _first_tensor_arg(node):
if _is_optional_tensor(arg):
continue
if _is_vector_tensor(arg):
return f"{arg.spelling}.at(0)"
return arg
if "Tensor" in arg.type.spelling:
return arg.spelling
return arg
return None

def _default_impl_index_expr(node):
first_tensor = _first_tensor_arg(node)
if first_tensor is None:
def _unique_local_name(node, base):
argument_names = {arg.spelling for arg in node.get_arguments()}
name = base
while name in argument_names:
name += "_"
return name

def _tensor_conversion_expr(arg):
if _is_vector_tensor(arg):
return f"VectorTensorFromPybind11Handle({arg.spelling})"
return f"TensorFromPybind11Handle({arg.spelling})"

def _default_impl_index_expr(node, converted_first_tensor_name=None):
first_tensor_arg = _first_tensor_arg(node)
if first_tensor_arg is None:
return "0"
return (
f"DefaultImplementationIndexFor{symbol_name}("
f"DeviceFromPybind11Handle({first_tensor}).type())"
)

if converted_first_tensor_name is not None:
first_tensor = converted_first_tensor_name
if _is_vector_tensor(first_tensor_arg):
first_tensor += ".at(0)"
device_type = f"{first_tensor}.device().type()"
else:
first_tensor = first_tensor_arg.spelling
if _is_vector_tensor(first_tensor_arg):
first_tensor += ".at(0)"
device_type = f"DeviceFromPybind11Handle({first_tensor}).type()"

return f"DefaultImplementationIndexFor{symbol_name}({device_type})"

def _generate_init(constructor):
constructor_params = _generate_params(constructor)
Expand Down Expand Up @@ -664,6 +692,20 @@ def _generate_call(op_name, call, method=True):
call_args = _generate_arguments(call)

if not method:
first_tensor_arg = _first_tensor_arg(call)
converted_first_tensor_name = None
first_tensor_conversion = ""
if first_tensor_arg is not None:
converted_first_tensor_name = _unique_local_name(
call, "converted_first_tensor"
)
first_tensor_conversion = (
f" auto {converted_first_tensor_name}"
f"{{{_tensor_conversion_expr(first_tensor_arg)}}};\n"
)
call_args = _generate_arguments(
call, first_tensor_arg, converted_first_tensor_name
)
params = (
f"{call_params}, std::uintptr_t stream, "
"std::optional<std::size_t> implementation_index"
Expand All @@ -673,17 +715,24 @@ def _generate_call(op_name, call, method=True):
)
py_args = _generate_py_args(call)
py_args_str = f"{py_args}, " if py_args else ""
default_impl_index = _default_impl_index_expr(call)
default_impl_index = _default_impl_index_expr(
call, converted_first_tensor_name
)

return (
f' m.def("{op_name}", []({params}) {{\n'
f" Handle handle;\n"
f" if (stream) {{\n"
f" handle.set_stream(reinterpret_cast<void*>(stream));\n"
f" }}\n"
f"{first_tensor_conversion}"
f" Config config;\n"
f" config.set_implementation_index(\n"
f" implementation_index.value_or({default_impl_index}));\n"
f" if (implementation_index.has_value()) {{\n"
f" config.set_implementation_index(*implementation_index);\n"
f" }} else {{\n"
f" config.set_implementation_index(\n"
f" {default_impl_index});\n"
f" }}\n"
f" return generated_dispatch::Call{symbol_name}(handle, config, {call_args});\n"
f' }}, {py_args_str}py::kw_only(), py::arg("stream") = 0, py::arg("implementation_index") = py::none());'
)
Expand Down Expand Up @@ -741,6 +790,7 @@ def _overload_order_key(node):

#include <pybind11/pybind11.h>
#include <pybind11/stl.h>
#include <utility>

#include "base/{op_name}.h"
#include "config.h"
Expand Down
49 changes: 48 additions & 1 deletion tests/test_generate_wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,13 +234,60 @@ class Mul {
"DefaultImplementationIndexForMul(DeviceFromPybind11Handle(input).type()))"
) in text
assert "std::optional<std::size_t> implementation_index" in text
assert "if (implementation_index.has_value())" in text
assert "config.set_implementation_index(*implementation_index)" in text
assert "auto converted_first_tensor{TensorFromPybind11Handle(input)};" in text
assert (
"implementation_index.value_or("
"DefaultImplementationIndexForMul(converted_first_tensor.device().type()))"
) in text
assert "std::move(converted_first_tensor)" in text
assert text.count("DeviceFromPybind11Handle(input)") == 1
assert (
"config.set_implementation_index("
"DefaultImplementationIndexForMul(DeviceFromPybind11Handle(input).type()))"
) in text
assert "implementation_index.value_or(" not in text
assert 'py::arg("implementation_index") = py::none()' in text


def test_pybind_default_implementation_reuses_first_vector_tensor(
monkeypatch, tmp_path
):
module = _load_generator_module()
base_header = tmp_path / "cat.h"
base_header.write_text(
"""
class Cat {
public:
virtual void operator()(const std::vector<Tensor> inputs, Tensor out) const = 0;
};
"""
)
monkeypatch.setattr(module, "_find_base_header", lambda op_name: base_header)

operator = module._Operator(
"cat",
constructors=[],
calls=[
module._ParsedFunction(
[
module._ParsedArgument("const std::vector<Tensor>", "inputs"),
module._ParsedArgument("Tensor", "out"),
]
)
],
)

text = module._generate_pybind11(operator)

assert (
"auto converted_first_tensor{VectorTensorFromPybind11Handle(inputs)};" in text
)
assert "converted_first_tensor.at(0).device().type()" in text
assert "std::move(converted_first_tensor)" in text
assert "DeviceFromPybind11Handle(inputs.at(0))" not in text


_DTYPE_OP_SOURCE = """
namespace infini::ops {

Expand Down
Loading