Add JavaScript bindings and WASM support
Co-authored-by: Matija Kecman <matijak@google.com> Co-authored-by: Sebastian Noreña Rendón <sebas.norena@creativa77.com.ar> Co-authored-by: Kyle Bayes <kylebayes@google.com> PiperOrigin-RevId: 826479070 Change-Id: I25acf36e6c20b091d8492e10798d028ae66f6733
This commit is contained in:
committed by
Copybara-Service
parent
57f7145806
commit
4086261714
@@ -0,0 +1,34 @@
|
||||
# Copyright 2025 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Generator for the constants."""
|
||||
|
||||
from wasm.codegen.helpers import common
|
||||
|
||||
|
||||
# TODO(manevi): Delete this file and use the genrule to handle the file copying
|
||||
class Generator:
|
||||
"""Generator for the constants."""
|
||||
|
||||
def run(self):
|
||||
"""Runs the generator."""
|
||||
template_cc_file, output_cc_file = common.get_file_path(
|
||||
"templates", "generated", "constants.cc"
|
||||
)
|
||||
|
||||
with open(template_cc_file, "r") as f_template:
|
||||
template_content = f_template.read()
|
||||
|
||||
with open(output_cc_file, "w") as f_output:
|
||||
f_output.write(template_content)
|
||||
@@ -0,0 +1,47 @@
|
||||
# Copyright 2025 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Generates Embind bindings for MuJoCo enums."""
|
||||
|
||||
from typing import Mapping
|
||||
|
||||
from introspect import ast_nodes
|
||||
|
||||
from wasm.codegen.helpers import code_builder
|
||||
|
||||
|
||||
class Generator:
|
||||
"""Generates Embind code for MuJoCo enums."""
|
||||
|
||||
def __init__(self, enums: Mapping[str, ast_nodes.EnumDecl]):
|
||||
self.enums = enums
|
||||
|
||||
def _generate_enum_binding(self, enum: ast_nodes.EnumDecl) -> str:
|
||||
"""Generates the Embind code for a single enum."""
|
||||
|
||||
code = f'{code_builder.INDENT}enum_<{enum.name}>("{enum.name}")'
|
||||
|
||||
for value_name in enum.values:
|
||||
code += f'\n{2*code_builder.INDENT}.value("{value_name}", {value_name})'
|
||||
|
||||
code += ";"
|
||||
return code
|
||||
|
||||
def generate(self) -> str:
|
||||
"""Generates all Embind code for the provided enums."""
|
||||
|
||||
code = []
|
||||
for enum in self.enums.values():
|
||||
code.append(self._generate_enum_binding(enum))
|
||||
return "\n\n".join(code) + "\n"
|
||||
@@ -0,0 +1,62 @@
|
||||
# Copyright 2025 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from absl.testing import absltest
|
||||
|
||||
from introspect import ast_nodes
|
||||
|
||||
from wasm.codegen.generators import enums
|
||||
|
||||
|
||||
class EnumsGeneratorTest(absltest.TestCase):
|
||||
|
||||
def test_generate_enum_bindings(self):
|
||||
|
||||
generator = enums.Generator({
|
||||
"TestEnum": ast_nodes.EnumDecl(
|
||||
name="TestEnum",
|
||||
declname="enum TestEnum_",
|
||||
values={"FIRST_VAL": 0, "SECOND_VAL": 1, "THIRD_VAL": 2},
|
||||
),
|
||||
"AnotherEnum": ast_nodes.EnumDecl(
|
||||
name="AnotherEnum",
|
||||
declname="enum AnotherEnum_",
|
||||
values={"ALPHA": 100, "BETA": 200},
|
||||
),
|
||||
"EmptyEnum": ast_nodes.EnumDecl(
|
||||
name="EmptyEnum",
|
||||
declname="enum EmptyEnum_",
|
||||
values={},
|
||||
),
|
||||
})
|
||||
|
||||
expected_code = """ enum_<TestEnum>("TestEnum")
|
||||
.value("FIRST_VAL", FIRST_VAL)
|
||||
.value("SECOND_VAL", SECOND_VAL)
|
||||
.value("THIRD_VAL", THIRD_VAL);
|
||||
|
||||
enum_<AnotherEnum>("AnotherEnum")
|
||||
.value("ALPHA", ALPHA)
|
||||
.value("BETA", BETA);
|
||||
|
||||
enum_<EmptyEnum>("EmptyEnum");
|
||||
"""
|
||||
|
||||
actual_code = generator.generate()
|
||||
|
||||
self.assertEqual(actual_code, expected_code)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
absltest.main()
|
||||
@@ -0,0 +1,93 @@
|
||||
# Copyright 2025 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Generates Embind bindings for MuJoCo functions."""
|
||||
|
||||
import pathlib
|
||||
from typing import List, Mapping, TypeAlias
|
||||
|
||||
from introspect import ast_nodes
|
||||
|
||||
from wasm.codegen.helpers import code_builder
|
||||
from wasm.codegen.helpers import function_utils
|
||||
|
||||
FunctionDecl: TypeAlias = ast_nodes.FunctionDecl
|
||||
FunctionParameterDecl: TypeAlias = ast_nodes.FunctionParameterDecl
|
||||
PointerType: TypeAlias = ast_nodes.PointerType
|
||||
ValueType: TypeAlias = ast_nodes.ValueType
|
||||
Path: TypeAlias = pathlib.Path
|
||||
|
||||
|
||||
class Generator:
|
||||
"""Generates Embind bindings for MuJoCo functions."""
|
||||
|
||||
def __init__(self, functions: Mapping[str, FunctionDecl]):
|
||||
self.direct_bind_functions: List[FunctionDecl] = []
|
||||
self.wrapper_bind_functions: List[FunctionDecl] = []
|
||||
|
||||
for func in functions.values():
|
||||
if function_utils.should_be_wrapped(func):
|
||||
self.wrapper_bind_functions.append(func)
|
||||
else:
|
||||
self.direct_bind_functions.append(func)
|
||||
|
||||
def _generate_wrappers(self) -> str:
|
||||
"""Generates Embind bindings for all functions that need wrappers."""
|
||||
|
||||
code = []
|
||||
for func in self.wrapper_bind_functions:
|
||||
wrapper_code = function_utils.generate_function_wrapper(func)
|
||||
code.append(wrapper_code)
|
||||
|
||||
return "\n\n".join(code)
|
||||
|
||||
def _generate_direct_bindable_functions(self) -> str:
|
||||
"""Generates Embind bindings for all directly bindable functions."""
|
||||
|
||||
result = ""
|
||||
for func in self.direct_bind_functions:
|
||||
result += code_builder.INDENT
|
||||
result += self._generate_function_binding(func)
|
||||
|
||||
return result
|
||||
|
||||
def _generate_function_binding(
|
||||
self, func: FunctionDecl, is_wrapper=False
|
||||
) -> str:
|
||||
"""Generates the Embind code for a single function."""
|
||||
|
||||
js_name, cpp_func = func.name, func.name
|
||||
if is_wrapper:
|
||||
cpp_func += "_wrapper"
|
||||
|
||||
return f'function("{js_name}", &{cpp_func});\n'
|
||||
|
||||
def _generate_wrapper_bindable_functions(self) -> str:
|
||||
"""Generates Embind bindings for all functions that need wrappers."""
|
||||
|
||||
result = ""
|
||||
for func in self.wrapper_bind_functions:
|
||||
result += code_builder.INDENT
|
||||
result += self._generate_function_binding(func, True)
|
||||
|
||||
return result
|
||||
|
||||
def generate(self) -> tuple[str, str]:
|
||||
"""Generates the bindings file for all functions."""
|
||||
|
||||
wrapper_functions = self._generate_wrappers()
|
||||
function_bindings = self._generate_direct_bindable_functions()
|
||||
function_bindings += self._generate_wrapper_bindable_functions()
|
||||
|
||||
return wrapper_functions, function_bindings
|
||||
@@ -0,0 +1,67 @@
|
||||
# Copyright 2025 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from absl.testing import absltest
|
||||
|
||||
from introspect import ast_nodes
|
||||
|
||||
from wasm.codegen.generators import functions
|
||||
|
||||
|
||||
class FunctionsGeneratorTest(absltest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.generator = functions.Generator({})
|
||||
self.int_type = ast_nodes.ValueType(name="int")
|
||||
|
||||
def test_generate_function_binding_simple_case(self):
|
||||
func_simple_void = ast_nodes.FunctionDecl(
|
||||
name="do_nothing",
|
||||
return_type=ast_nodes.ValueType(name="void"),
|
||||
parameters=tuple(),
|
||||
doc="doc",
|
||||
)
|
||||
self.assertEqual(
|
||||
self.generator._generate_function_binding(func_simple_void),
|
||||
'function("do_nothing", &do_nothing);\n',
|
||||
)
|
||||
|
||||
def test_generate_direct_bindable_functions_simple_filter(self):
|
||||
direct_bind = ast_nodes.FunctionDecl(
|
||||
name="direct_bind",
|
||||
return_type=self.int_type,
|
||||
parameters=(
|
||||
ast_nodes.FunctionParameterDecl(name="val", type=self.int_type),
|
||||
),
|
||||
doc="doc",
|
||||
)
|
||||
needs_wrap = ast_nodes.FunctionDecl(
|
||||
name="needs_wrap",
|
||||
return_type=ast_nodes.PointerType(inner_type=self.int_type),
|
||||
parameters=tuple(),
|
||||
doc="doc",
|
||||
)
|
||||
self.generator = functions.Generator({
|
||||
"direct1": direct_bind,
|
||||
"wrapped1": needs_wrap,
|
||||
})
|
||||
|
||||
generated_code = self.generator._generate_direct_bindable_functions()
|
||||
self.assertIn('function("direct_bind", &direct_bind);\n', generated_code)
|
||||
self.assertNotIn('function("needs_wrap", &needs_wrap);\n', generated_code)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
absltest.main()
|
||||
@@ -0,0 +1,88 @@
|
||||
# Copyright 2025 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Generates Embind bindings for MuJoCo structs."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from wasm.codegen.helpers import constants
|
||||
from wasm.codegen.helpers import structs_parser
|
||||
from wasm.codegen.helpers import structs_wrappers_data
|
||||
|
||||
|
||||
class Generator:
|
||||
"""Generates C++ code for binding and wrapping MuJoCo structs."""
|
||||
|
||||
def __init__(self):
|
||||
# Set up the correct input dict based on the structs we want to bind
|
||||
# and already have a wrapper manually created in the template/bindings.cc
|
||||
wrapped_structs = structs_wrappers_data.create_wrapped_structs_set_up_data(
|
||||
constants.STRUCTS_TO_BIND
|
||||
)
|
||||
|
||||
# Traverse the introspect dictionary to get the field
|
||||
# wrapper/bindings statements set up for each struct
|
||||
self.structs_to_bind_data = structs_parser.generate_wasm_bindings(
|
||||
wrapped_structs
|
||||
)
|
||||
|
||||
def generate_header(
|
||||
self
|
||||
) -> list[tuple[str, list[Optional[str]]]]:
|
||||
"""Generates C++ header file for binding and wrapping MuJoCo structs."""
|
||||
autogenned_struct_definitions = []
|
||||
markers_and_content = []
|
||||
|
||||
# Sort by struct name by dependency to ensure deterministic output order
|
||||
sorted_struct_names = structs_parser.sort_structs_by_dependency(
|
||||
constants.STRUCTS_TO_BIND
|
||||
)
|
||||
|
||||
for struct_name in sorted_struct_names:
|
||||
struct_data = self.structs_to_bind_data[struct_name]
|
||||
if struct_data.wrapped_header:
|
||||
autogenned_struct_definitions.append(
|
||||
struct_data.wrapped_header + "\n"
|
||||
)
|
||||
else:
|
||||
markers_and_content.append((
|
||||
f"// INSERT-GENERATED-{struct_data.wrap_name}-DEFINITIONS",
|
||||
[
|
||||
l.definition if l.definition else ""
|
||||
for l in struct_data.wrapped_fields
|
||||
],
|
||||
))
|
||||
markers_and_content.append((
|
||||
"// {{ AUTOGENNED_STRUCT_DEFINITIONS }}",
|
||||
autogenned_struct_definitions,
|
||||
))
|
||||
return markers_and_content
|
||||
|
||||
def generate_source(self) -> list[tuple[str, list[str]]]:
|
||||
"""Generates C++ source file for binding and wrapping MuJoCo structs."""
|
||||
constructors = [
|
||||
(
|
||||
f"// INSERT-GENERATED-{struct_data.wrap_name}-CONSTRUCTOR",
|
||||
[struct_data.wrapped_source],
|
||||
)
|
||||
for _, struct_data in self.structs_to_bind_data.items()
|
||||
]
|
||||
properties = [
|
||||
(
|
||||
f"// INSERT-GENERATED-{struct_data.wrap_name}-BINDINGS",
|
||||
[l.binding for l in struct_data.wrapped_fields],
|
||||
)
|
||||
for _, struct_data in self.structs_to_bind_data.items()
|
||||
]
|
||||
return constructors + properties
|
||||
Reference in New Issue
Block a user