diff --git a/python/generated_ops.py b/python/generated_ops.py new file mode 100644 index 00000000..51bf0127 --- /dev/null +++ b/python/generated_ops.py @@ -0,0 +1,317 @@ +# Copyright 2026 The TensorFlow MUSA Authors. All Rights Reserved. +# +# 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. +# ============================================================================== + +"""Generated public wrappers for MUSA extension ops.""" + +from . import raw_ops + + +def batch_mat_mul_v2(x, y, adj_x=False, adj_y=False, name=None): + return raw_ops.musa_batch_mat_mul_v2( + x=x, + y=y, + adj_x=adj_x, + adj_y=adj_y, + name=name, + ) + + +def bias_add_relu_mat_mul(input, bias, other, relu_input_slot, transpose_a=False, transpose_b=False, name=None): + return raw_ops.musa_bias_add_relu_mat_mul( + input=input, + bias=bias, + other=other, + relu_input_slot=relu_input_slot, + transpose_a=transpose_a, + transpose_b=transpose_b, + name=name, + ) + + +def clip(x, lo, hi, name=None): + return raw_ops.musa_clip( + x=x, + lo=lo, + hi=hi, + name=name, + ) + + +def concat_mat_mul(inputs, axis, other, concat_input_idx, transpose_a=False, transpose_b=False, name=None): + return raw_ops.musa_concat_mat_mul( + inputs=inputs, + axis=axis, + other=other, + concat_input_idx=concat_input_idx, + transpose_a=transpose_a, + transpose_b=transpose_b, + name=name, + ) + + +def dropout(x, rate=0.5, seed=0, offset=0, name=None): + return raw_ops.musa_dropout( + x=x, + rate=rate, + seed=seed, + offset=offset, + name=name, + ) + + +def dropout_grad(grad, mask, rate=0.5, name=None): + return raw_ops.musa_dropout_grad( + grad=grad, + mask=mask, + rate=rate, + name=name, + ) + + +def gelu(x, approximate=False, name=None): + return raw_ops.musa_gelu(x=x, approximate=approximate, name=name) + + +def interact(input, name=None): + return raw_ops.musa_interact(input=input, name=name) + + +def layer_norm(x, gamma, beta, epsilon=0.00001, name=None): + return raw_ops.musa_layer_norm( + x=x, + gamma=gamma, + beta=beta, + epsilon=epsilon, + name=name, + ) + + +def linear_activation(a, b, bias, activation='relu', alpha=0.0, transpose_a=False, transpose_b=False, name=None): + return raw_ops.musa_linear_activation( + a=a, + b=b, + bias=bias, + activation=activation, + alpha=alpha, + transpose_a=transpose_a, + transpose_b=transpose_b, + name=name, + ) + + +def mat_mul(a, b, transpose_a=False, transpose_b=False, name=None): + return raw_ops.musa_mat_mul( + a=a, + b=b, + transpose_a=transpose_a, + transpose_b=transpose_b, + name=name, + ) + + +def matmul_bias_add(a, b, bias, transpose_a=False, transpose_b=False, name=None): + return raw_ops.musa_mat_mul_bias_add( + a=a, + b=b, + bias=bias, + transpose_a=transpose_a, + transpose_b=transpose_b, + name=name, + ) + + +def maximum(x, y, name=None): + return raw_ops.musa_maximum(x=x, y=y, name=name) + + +def mean(input, reduction_indices, keep_dims=False, name=None): + return raw_ops.musa_mean( + input=input, + reduction_indices=reduction_indices, + keep_dims=keep_dims, + name=name, + ) + + +def normalize(x, gamma, beta, epsilon=1e-11, max_std=float('inf'), name=None): + return raw_ops.musa_normalize( + x=x, + gamma=gamma, + beta=beta, + epsilon=epsilon, + max_std=max_std, + name=name, + ) + + +def pln_cascade(norm_out, adpos, add_input, bias_input, use_table=False, table_index=0, select_on_true=True, name=None): + return raw_ops.musa_pln_cascade( + norm_out=norm_out, + adpos=adpos, + add_input=add_input, + bias_input=bias_input, + use_table=use_table, + table_index=table_index, + select_on_true=select_on_true, + name=name, + ) + + +def pln_cascade_block(norm_out, add_input, bias_input, gates, table_indices, select_on_true, name=None): + return raw_ops.musa_pln_cascade_block( + norm_out=norm_out, + add_input=add_input, + bias_input=bias_input, + gates=gates, + table_indices=table_indices, + select_on_true=select_on_true, + name=name, + ) + + +def prelu(x, alpha, name=None): + return raw_ops.musa_p_relu(x=x, alpha=alpha, name=name) + + +def reshape_mat_mul(x, w, transpose_b=False, name=None): + return raw_ops.musa_reshape_mat_mul( + x=x, + w=w, + transpose_b=transpose_b, + name=name, + ) + + +def resource_apply_adam_mixed(var, m, v, beta1_power, beta2_power, lr, beta1, beta2, epsilon, grad, use_locking=False, use_nesterov=False, name=None): + return raw_ops.musa_resource_apply_adam_mixed( + var=var, + m=m, + v=v, + beta1_power=beta1_power, + beta2_power=beta2_power, + lr=lr, + beta1=beta1, + beta2=beta2, + epsilon=epsilon, + grad=grad, + use_locking=use_locking, + use_nesterov=use_nesterov, + name=name, + ) + + +def resource_sparse_apply_adam(var, m, v, beta1_power, beta2_power, lr, beta1, beta2, epsilon, grad, indices, use_locking=False, name=None): + return raw_ops.musa_resource_sparse_apply_adam( + var=var, + m=m, + v=v, + beta1_power=beta1_power, + beta2_power=beta2_power, + lr=lr, + beta1=beta1, + beta2=beta2, + epsilon=epsilon, + grad=grad, + indices=indices, + use_locking=use_locking, + name=name, + ) + + +def shifted_affine_map(data_left, mask, sliced_var_right, name=None): + return raw_ops.musa_shifted_affine_map( + data_left=data_left, + mask=mask, + sliced_var_right=sliced_var_right, + name=name, + ) + + +def tensor_dot(a, b, axes_a, axes_b, name=None): + return raw_ops.musa_tensor_dot( + a=a, + b=b, + axes_a=axes_a, + axes_b=axes_b, + name=name, + ) + + +def tensor_dot_bias(a, b, bias, axes_a, axes_b, name=None): + return raw_ops.musa_tensor_dot_bias( + a=a, + b=b, + bias=bias, + axes_a=axes_a, + axes_b=axes_b, + name=name, + ) + + +def token_mixer(x, num_T, num_H, d_k, name=None): + return raw_ops.musa_token_mixer( + x=x, + num_T=num_T, + num_H=num_H, + d_k=d_k, + name=name, + ) + + +def resource_apply_nadam(var, m, v, beta1_power, beta2_power, lr, beta1, beta2, epsilon, grad, use_locking=False, name=None): + return raw_ops.ResourceApplyNadam( + var=var, + m=m, + v=v, + beta1_power=beta1_power, + beta2_power=beta2_power, + lr=lr, + beta1=beta1, + beta2=beta2, + epsilon=epsilon, + grad=grad, + use_locking=use_locking, + name=name, + ) + + +__all__ = [ + "batch_mat_mul_v2", + "bias_add_relu_mat_mul", + "clip", + "concat_mat_mul", + "dropout", + "dropout_grad", + "gelu", + "interact", + "layer_norm", + "linear_activation", + "mat_mul", + "matmul_bias_add", + "maximum", + "mean", + "normalize", + "pln_cascade", + "pln_cascade_block", + "prelu", + "reshape_mat_mul", + "resource_apply_adam_mixed", + "resource_apply_nadam", + "resource_sparse_apply_adam", + "shifted_affine_map", + "tensor_dot", + "tensor_dot_bias", + "token_mixer", +] diff --git a/python/op_manifest.py b/python/op_manifest.py new file mode 100644 index 00000000..8714fd10 --- /dev/null +++ b/python/op_manifest.py @@ -0,0 +1,125 @@ +# Copyright 2026 The TensorFlow MUSA Authors. All Rights Reserved. +# +# 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. +# ============================================================================== + +"""Manifest of custom MUSA TensorFlow ops exposed to Python.""" + +CUSTOM_OPS = [{'op': 'MusaBatchMatMulV2', + 'raw': 'musa_batch_mat_mul_v2', + 'api': 'batch_mat_mul_v2', + 'source': 'musa_ext/kernels/math/musa_matmul_op.cc'}, + {'op': 'MusaBiasAddReluMatMul', + 'raw': 'musa_bias_add_relu_mat_mul', + 'api': 'bias_add_relu_mat_mul', + 'source': 'musa_ext/kernels/fusion/musa_rgprojection_fusion_op.cc'}, + {'op': 'MusaClip', + 'raw': 'musa_clip', + 'api': 'clip', + 'source': 'musa_ext/kernels/fusion/musa_clip_op.cc'}, + {'op': 'MusaConcatMatMul', + 'raw': 'musa_concat_mat_mul', + 'api': 'concat_mat_mul', + 'source': 'musa_ext/kernels/fusion/musa_concat_matmul_op.cc'}, + {'op': 'MusaDropout', + 'raw': 'musa_dropout', + 'api': 'dropout', + 'source': 'musa_ext/kernels/nn/musa_dropout_op.cc'}, + {'op': 'MusaDropoutGrad', + 'raw': 'musa_dropout_grad', + 'api': 'dropout_grad', + 'source': 'musa_ext/kernels/nn/musa_dropout_op.cc'}, + {'op': 'MusaGelu', + 'raw': 'musa_gelu', + 'api': 'gelu', + 'source': 'musa_ext/kernels/fusion/musa_gelu_op.cc'}, + {'op': 'MusaInteract', + 'raw': 'musa_interact', + 'api': 'interact', + 'source': 'musa_ext/kernels/array/musa_tensorinteraction_op.cc'}, + {'op': 'MusaLayerNorm', + 'raw': 'musa_layer_norm', + 'api': 'layer_norm', + 'source': 'musa_ext/kernels/fusion/musa_layernorm_op.cc'}, + {'op': 'MusaLinearActivation', + 'raw': 'musa_linear_activation', + 'api': 'linear_activation', + 'source': 'musa_ext/kernels/fusion/musa_linear_relu_op.cc'}, + {'op': 'MusaMatMul', + 'raw': 'musa_mat_mul', + 'api': 'mat_mul', + 'source': 'musa_ext/kernels/math/musa_matmul_op.cc'}, + {'op': 'MusaMatMulBiasAdd', + 'raw': 'musa_mat_mul_bias_add', + 'api': 'matmul_bias_add', + 'source': 'musa_ext/kernels/fusion/musa_matmul_bias_op.cc'}, + {'op': 'MusaMaximum', + 'raw': 'musa_maximum', + 'api': 'maximum', + 'source': 'musa_ext/kernels/math/musa_maximum_op.cc'}, + {'op': 'MusaMean', + 'raw': 'musa_mean', + 'api': 'mean', + 'source': 'musa_ext/kernels/math/musa_mean_op.cc'}, + {'op': 'MusaNormalize', + 'raw': 'musa_normalize', + 'api': 'normalize', + 'source': 'musa_ext/kernels/fusion/musa_normalize_fusion_op.cc'}, + {'op': 'MusaPlnCascade', + 'raw': 'musa_pln_cascade', + 'api': 'pln_cascade', + 'source': 'musa_ext/kernels/fusion/musa_pln_cascade_op.cc'}, + {'op': 'MusaPlnCascadeBlock', + 'raw': 'musa_pln_cascade_block', + 'api': 'pln_cascade_block', + 'source': 'musa_ext/kernels/fusion/musa_pln_cascade_block_op.cc'}, + {'op': 'MusaPRelu', + 'raw': 'musa_p_relu', + 'api': 'prelu', + 'source': 'musa_ext/kernels/fusion/musa_prelu_fusion_op.cc'}, + {'op': 'MusaReshapeMatMul', + 'raw': 'musa_reshape_mat_mul', + 'api': 'reshape_mat_mul', + 'source': 'musa_ext/kernels/fusion/musa_reshape_matmul_op.cc'}, + {'op': 'MusaResourceApplyAdamMixed', + 'raw': 'musa_resource_apply_adam_mixed', + 'api': 'resource_apply_adam_mixed', + 'source': 'musa_ext/kernels/training/musa_applyadam_mixed_op.cc'}, + {'op': 'MusaResourceSparseApplyAdam', + 'raw': 'musa_resource_sparse_apply_adam', + 'api': 'resource_sparse_apply_adam', + 'source': 'musa_ext/kernels/training/musa_apply_sparse_adam_op.cc'}, + {'op': 'MusaShiftedAffineMap', + 'raw': 'musa_shifted_affine_map', + 'api': 'shifted_affine_map', + 'source': 'musa_ext/kernels/fusion/musa_shifted_affine_map_op.cc'}, + {'op': 'MusaTensorDot', + 'raw': 'musa_tensor_dot', + 'api': 'tensor_dot', + 'source': 'musa_ext/kernels/fusion/musa_tensordot_op.cc'}, + {'op': 'MusaTensorDotBias', + 'raw': 'musa_tensor_dot_bias', + 'api': 'tensor_dot_bias', + 'source': 'musa_ext/kernels/fusion/musa_tensordot_bias_op.cc'}, + {'op': 'MusaTokenMixer', + 'raw': 'musa_token_mixer', + 'api': 'token_mixer', + 'source': 'musa_ext/kernels/fusion/musa_tokenmixer_op.cc'}, + {'op': 'ResourceApplyNadam', + 'raw': 'ResourceApplyNadam', + 'api': 'resource_apply_nadam', + 'source': 'musa_ext/kernels/training/musa_resource_apply_nadam_op.cc'}] + +CUSTOM_OP_NAMES = tuple(entry["op"] for entry in CUSTOM_OPS) +RAW_OP_NAMES = tuple(entry["raw"] for entry in CUSTOM_OPS) +PUBLIC_API_NAMES = tuple(entry["api"] for entry in CUSTOM_OPS) diff --git a/python/ops.py b/python/ops.py index e4140ca0..0c052598 100644 --- a/python/ops.py +++ b/python/ops.py @@ -13,157 +13,10 @@ # limitations under the License. # ============================================================================== -"""Convenience wrappers for selected MUSA extension ops.""" +"""Convenience wrappers for MUSA extension ops.""" from . import raw_ops +from .generated_ops import * +from .generated_ops import __all__ as _generated_all - -def clip(x, lo, hi, name=None): - return raw_ops.musa_clip(x=x, lo=lo, hi=hi, name=name) - - -def layer_norm(x, gamma, beta, epsilon=0.00001, name=None): - return raw_ops.musa_layer_norm( - x=x, - gamma=gamma, - beta=beta, - epsilon=epsilon, - name=name, - ) - - -def shifted_affine_map(data_left, mask, sliced_var_right, name=None): - return raw_ops.musa_shifted_affine_map( - data_left=data_left, - mask=mask, - sliced_var_right=sliced_var_right, - name=name, - ) - - -def interact(input, name=None): - return raw_ops.musa_interact(input=input, name=name) - - -def dropout(x, rate=0.5, seed=0, offset=0, name=None): - return raw_ops.musa_dropout( - x=x, - rate=rate, - seed=seed, - offset=offset, - name=name, - ) - - -def dropout_grad(grad, mask, rate=0.5, name=None): - return raw_ops.musa_dropout_grad( - grad=grad, - mask=mask, - rate=rate, - name=name, - ) - - -def resource_sparse_apply_adam( - var, - m, - v, - beta1_power, - beta2_power, - lr, - beta1, - beta2, - epsilon, - grad, - indices, - use_locking=False, - name=None, -): - return raw_ops.musa_resource_sparse_apply_adam( - var=var, - m=m, - v=v, - beta1_power=beta1_power, - beta2_power=beta2_power, - lr=lr, - beta1=beta1, - beta2=beta2, - epsilon=epsilon, - grad=grad, - indices=indices, - use_locking=use_locking, - name=name, - ) - - -def gelu(x, approximate=False, name=None): - return raw_ops.musa_gelu( - x=x, - approximate=approximate, - name=name, - ) - - -def reshape_mat_mul(x, w, transpose_b=False, name=None): - return raw_ops.musa_reshape_mat_mul( - x=x, - w=w, - transpose_b=transpose_b, - name=name, - ) - - -def matmul_bias_add(a, b, bias, transpose_a=False, transpose_b=False, name=None): - return raw_ops.musa_mat_mul_bias_add( - a=a, - b=b, - bias=bias, - transpose_a=transpose_a, - transpose_b=transpose_b, - name=name, - ) - - -def resource_apply_nadam( - var, - m, - v, - beta1_power, - beta2_power, - lr, - beta1, - beta2, - epsilon, - grad, - use_locking=False, - name=None, -): - return raw_ops.ResourceApplyNadam( - var=var, - m=m, - v=v, - beta1_power=beta1_power, - beta2_power=beta2_power, - lr=lr, - beta1=beta1, - beta2=beta2, - epsilon=epsilon, - grad=grad, - use_locking=use_locking, - name=name, - ) - - -__all__ = [ - "clip", - "dropout", - "dropout_grad", - "gelu", - "interact", - "layer_norm", - "matmul_bias_add", - "reshape_mat_mul", - "resource_apply_nadam", - "resource_sparse_apply_adam", - "shifted_affine_map", -] +__all__ = list(_generated_all) diff --git a/test/ops/python_api_op_test.py b/test/ops/python_api_op_test.py index d336a46c..989d5739 100644 --- a/test/ops/python_api_op_test.py +++ b/test/ops/python_api_op_test.py @@ -15,7 +15,9 @@ """Tests for the public MUSA Python op API.""" +import inspect import importlib.util +import re import sys import types import unittest @@ -26,8 +28,12 @@ import tensorflow as tf -if "tensorflow_musa" not in sys.modules: - package_dir = Path(__file__).resolve().parents[2] / "python" +package_dir = Path(__file__).resolve().parents[2] / "python" +module = sys.modules.get("tensorflow_musa") +if module is None or not str(getattr(module, "__file__", "")).startswith(str(package_dir)): + for module_name in list(sys.modules): + if module_name == "tensorflow_musa" or module_name.startswith("tensorflow_musa."): + del sys.modules[module_name] spec = importlib.util.spec_from_file_location( "tensorflow_musa", package_dir / "__init__.py", @@ -35,10 +41,15 @@ ) tensorflow_musa = importlib.util.module_from_spec(spec) sys.modules["tensorflow_musa"] = tensorflow_musa - spec.loader.exec_module(tensorflow_musa) + from tensorflow.python.framework import load_library as tf_load_library + + with mock.patch("tensorflow.load_op_library", return_value=types.SimpleNamespace()): + with mock.patch.object(tf_load_library, "load_pluggable_device_library"): + spec.loader.exec_module(tensorflow_musa) import tensorflow_musa from tensorflow_musa import ops, raw_ops +from tensorflow_musa.op_manifest import CUSTOM_OPS, PUBLIC_API_NAMES, RAW_OP_NAMES class PythonApiOpTest(unittest.TestCase): @@ -57,12 +68,58 @@ def _restore_raw_op(self, name, previous, sentinel): else: ops.raw_ops.__dict__[name] = previous + def _skip_unless_raw_op_available(self, name): + if not hasattr(tensorflow_musa.get_musa_ops(), name): + self.skipTest(f"MUSA raw op {name!r} is not available in this test environment") + def testPackageExportsOpModules(self): self.assertIs(tensorflow_musa.ops, ops) self.assertIs(tensorflow_musa.raw_ops, raw_ops) self.assertIn("ops", tensorflow_musa.__all__) self.assertIn("raw_ops", tensorflow_musa.__all__) + def testAllCustomOpsAreInPublicAll(self): + self.assertEqual(set(PUBLIC_API_NAMES), set(ops.__all__)) + for api_name in PUBLIC_API_NAMES: + with self.subTest(api_name=api_name): + self.assertTrue(callable(getattr(ops, api_name))) + + def testManifestCoversRegisteredCustomOps(self): + root = Path(__file__).resolve().parents[2] + pattern = re.compile(r'REGISTER_OP\("([^"]+)"\)') + registered_ops = set() + for source in (root / "musa_ext" / "kernels").rglob("*.cc"): + registered_ops.update(pattern.findall(source.read_text(encoding="utf-8"))) + + self.assertEqual(registered_ops, {entry["op"] for entry in CUSTOM_OPS}) + + def testAllGeneratedWrappersDelegateToRawOps(self): + for entry in CUSTOM_OPS: + with self.subTest(api=entry["api"]): + wrapper = getattr(ops, entry["api"]) + signature = inspect.signature(wrapper) + kwargs = {} + expected_kwargs = {} + for name, parameter in signature.parameters.items(): + if parameter.default is inspect.Parameter.empty: + value = f"{entry['api']}_{name}" + elif name == "name": + value = f"{entry['api']}_name" + else: + value = parameter.default + kwargs[name] = value + expected_kwargs[name] = value + + op = self._patch_raw_op(entry["raw"], "result") + self.assertEqual(wrapper(**kwargs), "result") + op.assert_called_once_with(**expected_kwargs) + + def testRawOpManifestNamesAreUnique(self): + self.assertEqual(len(CUSTOM_OPS), len(PUBLIC_API_NAMES)) + self.assertEqual(len(CUSTOM_OPS), len(RAW_OP_NAMES)) + self.assertEqual(len(PUBLIC_API_NAMES), len(set(PUBLIC_API_NAMES))) + self.assertEqual(len(RAW_OP_NAMES), len(set(RAW_OP_NAMES))) + def testRawOpsDelegatesToGeneratedModule(self): generated = types.SimpleNamespace(musa_clip=lambda **kwargs: kwargs) @@ -254,11 +311,13 @@ def testResourceSparseApplyAdamWrapperDelegatesToRawOp(self): ) def testRealClipComputesExpectedValues(self): + self._skip_unless_raw_op_available("musa_clip") result = ops.clip(tf.constant([-2.0, 0.5, 3.0]), 0.0, 1.0) np.testing.assert_allclose(result.numpy(), [0.0, 0.5, 1.0]) def testRealLayerNormComputesExpectedValues(self): + self._skip_unless_raw_op_available("musa_layer_norm") x = tf.constant([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]) result = ops.layer_norm(x, tf.ones([3]), tf.zeros([3])) @@ -273,6 +332,7 @@ def testRealLayerNormComputesExpectedValues(self): np.testing.assert_allclose(result.numpy(), expected.numpy(), rtol=1e-5) def testRealShiftedAffineMapComputesExpectedValues(self): + self._skip_unless_raw_op_available("musa_shifted_affine_map") result = ops.shifted_affine_map( tf.ones([2, 3]), tf.ones([2, 3]), @@ -282,17 +342,20 @@ def testRealShiftedAffineMapComputesExpectedValues(self): np.testing.assert_allclose(result.numpy(), np.full([2, 3], 2.0)) def testRealInteractComputesExpectedValues(self): + self._skip_unless_raw_op_available("musa_interact") result = ops.interact(tf.ones([2, 3, 4])) np.testing.assert_allclose(result.numpy(), np.full([2, 3, 3], 4.0)) def testRealDropoutComputesOutputAndMask(self): + self._skip_unless_raw_op_available("musa_dropout") y, mask = ops.dropout(tf.ones([2, 3]), rate=0.5, seed=1, offset=0) self.assertEqual(y.shape, [2, 3]) self.assertEqual(mask.shape, [2, 3]) np.testing.assert_allclose(y.numpy(), mask.numpy().astype(np.float32) * 2.0) def testRealGeluComputesExpectedValues(self): + self._skip_unless_raw_op_available("musa_gelu") x = tf.constant([-2.0, -0.5, 0.0, 0.5, 2.0]) for approximate in [False, True]: with self.subTest(approximate=approximate): @@ -307,6 +370,7 @@ def testRealGeluComputesExpectedValues(self): ) def testRealReshapeMatMulComputesExpectedValues(self): + self._skip_unless_raw_op_available("musa_reshape_mat_mul") x = tf.constant([[[1.0, 2.0], [3.0, 4.0]]]) w = tf.constant([[5.0, 6.0], [7.0, 8.0]]) result = ops.reshape_mat_mul(x, w) @@ -314,6 +378,7 @@ def testRealReshapeMatMulComputesExpectedValues(self): np.testing.assert_allclose(result.numpy(), tf.matmul(x, w).numpy()) def testRealMatmulBiasAddComputesExpectedValues(self): + self._skip_unless_raw_op_available("musa_mat_mul_bias_add") a = tf.constant([[1.0, 2.0], [3.0, 4.0]]) b = tf.constant([[5.0, 6.0], [7.0, 8.0]]) bias = tf.constant([0.5, -0.5]) diff --git a/tools/generate_python_ops.py b/tools/generate_python_ops.py new file mode 100644 index 00000000..3b74eb69 --- /dev/null +++ b/tools/generate_python_ops.py @@ -0,0 +1,306 @@ +#!/usr/bin/env python3 +"""Generate Python wrappers for custom MUSA TensorFlow ops.""" + +from __future__ import annotations + +import ast +import pprint +import re +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +KERNELS_ROOT = ROOT / "musa_ext" / "kernels" +GENERATED_OPS = ROOT / "python" / "generated_ops.py" +OP_MANIFEST = ROOT / "python" / "op_manifest.py" + +OP_DEFINITIONS = [ + { + "op": "MusaBatchMatMulV2", + "raw": "musa_batch_mat_mul_v2", + "api": "batch_mat_mul_v2", + "source": "musa_ext/kernels/math/musa_matmul_op.cc", + "args": ["x", "y", "adj_x=False", "adj_y=False", "name=None"], + "call": ["x=x", "y=y", "adj_x=adj_x", "adj_y=adj_y", "name=name"], + }, + { + "op": "MusaBiasAddReluMatMul", + "raw": "musa_bias_add_relu_mat_mul", + "api": "bias_add_relu_mat_mul", + "source": "musa_ext/kernels/fusion/musa_rgprojection_fusion_op.cc", + "args": ["input", "bias", "other", "relu_input_slot", "transpose_a=False", "transpose_b=False", "name=None"], + "call": ["input=input", "bias=bias", "other=other", "relu_input_slot=relu_input_slot", "transpose_a=transpose_a", "transpose_b=transpose_b", "name=name"], + }, + { + "op": "MusaClip", + "raw": "musa_clip", + "api": "clip", + "source": "musa_ext/kernels/fusion/musa_clip_op.cc", + "args": ["x", "lo", "hi", "name=None"], + "call": ["x=x", "lo=lo", "hi=hi", "name=name"], + }, + { + "op": "MusaConcatMatMul", + "raw": "musa_concat_mat_mul", + "api": "concat_mat_mul", + "source": "musa_ext/kernels/fusion/musa_concat_matmul_op.cc", + "args": ["inputs", "axis", "other", "concat_input_idx", "transpose_a=False", "transpose_b=False", "name=None"], + "call": ["inputs=inputs", "axis=axis", "other=other", "concat_input_idx=concat_input_idx", "transpose_a=transpose_a", "transpose_b=transpose_b", "name=name"], + }, + { + "op": "MusaDropout", + "raw": "musa_dropout", + "api": "dropout", + "source": "musa_ext/kernels/nn/musa_dropout_op.cc", + "args": ["x", "rate=0.5", "seed=0", "offset=0", "name=None"], + "call": ["x=x", "rate=rate", "seed=seed", "offset=offset", "name=name"], + }, + { + "op": "MusaDropoutGrad", + "raw": "musa_dropout_grad", + "api": "dropout_grad", + "source": "musa_ext/kernels/nn/musa_dropout_op.cc", + "args": ["grad", "mask", "rate=0.5", "name=None"], + "call": ["grad=grad", "mask=mask", "rate=rate", "name=name"], + }, + { + "op": "MusaGelu", + "raw": "musa_gelu", + "api": "gelu", + "source": "musa_ext/kernels/fusion/musa_gelu_op.cc", + "args": ["x", "approximate=False", "name=None"], + "call": ["x=x", "approximate=approximate", "name=name"], + }, + { + "op": "MusaInteract", + "raw": "musa_interact", + "api": "interact", + "source": "musa_ext/kernels/array/musa_tensorinteraction_op.cc", + "args": ["input", "name=None"], + "call": ["input=input", "name=name"], + }, + { + "op": "MusaLayerNorm", + "raw": "musa_layer_norm", + "api": "layer_norm", + "source": "musa_ext/kernels/fusion/musa_layernorm_op.cc", + "args": ["x", "gamma", "beta", "epsilon=0.00001", "name=None"], + "call": ["x=x", "gamma=gamma", "beta=beta", "epsilon=epsilon", "name=name"], + }, + { + "op": "MusaLinearActivation", + "raw": "musa_linear_activation", + "api": "linear_activation", + "source": "musa_ext/kernels/fusion/musa_linear_relu_op.cc", + "args": ["a", "b", "bias", "activation='relu'", "alpha=0.0", "transpose_a=False", "transpose_b=False", "name=None"], + "call": ["a=a", "b=b", "bias=bias", "activation=activation", "alpha=alpha", "transpose_a=transpose_a", "transpose_b=transpose_b", "name=name"], + }, + { + "op": "MusaMatMul", + "raw": "musa_mat_mul", + "api": "mat_mul", + "source": "musa_ext/kernels/math/musa_matmul_op.cc", + "args": ["a", "b", "transpose_a=False", "transpose_b=False", "name=None"], + "call": ["a=a", "b=b", "transpose_a=transpose_a", "transpose_b=transpose_b", "name=name"], + }, + { + "op": "MusaMatMulBiasAdd", + "raw": "musa_mat_mul_bias_add", + "api": "matmul_bias_add", + "source": "musa_ext/kernels/fusion/musa_matmul_bias_op.cc", + "args": ["a", "b", "bias", "transpose_a=False", "transpose_b=False", "name=None"], + "call": ["a=a", "b=b", "bias=bias", "transpose_a=transpose_a", "transpose_b=transpose_b", "name=name"], + }, + { + "op": "MusaMaximum", + "raw": "musa_maximum", + "api": "maximum", + "source": "musa_ext/kernels/math/musa_maximum_op.cc", + "args": ["x", "y", "name=None"], + "call": ["x=x", "y=y", "name=name"], + }, + { + "op": "MusaMean", + "raw": "musa_mean", + "api": "mean", + "source": "musa_ext/kernels/math/musa_mean_op.cc", + "args": ["input", "reduction_indices", "keep_dims=False", "name=None"], + "call": ["input=input", "reduction_indices=reduction_indices", "keep_dims=keep_dims", "name=name"], + }, + { + "op": "MusaNormalize", + "raw": "musa_normalize", + "api": "normalize", + "source": "musa_ext/kernels/fusion/musa_normalize_fusion_op.cc", + "args": ["x", "gamma", "beta", "epsilon=1e-11", "max_std=float('inf')", "name=None"], + "call": ["x=x", "gamma=gamma", "beta=beta", "epsilon=epsilon", "max_std=max_std", "name=name"], + }, + { + "op": "MusaPlnCascade", + "raw": "musa_pln_cascade", + "api": "pln_cascade", + "source": "musa_ext/kernels/fusion/musa_pln_cascade_op.cc", + "args": ["norm_out", "adpos", "add_input", "bias_input", "use_table=False", "table_index=0", "select_on_true=True", "name=None"], + "call": ["norm_out=norm_out", "adpos=adpos", "add_input=add_input", "bias_input=bias_input", "use_table=use_table", "table_index=table_index", "select_on_true=select_on_true", "name=name"], + }, + { + "op": "MusaPlnCascadeBlock", + "raw": "musa_pln_cascade_block", + "api": "pln_cascade_block", + "source": "musa_ext/kernels/fusion/musa_pln_cascade_block_op.cc", + "args": ["norm_out", "add_input", "bias_input", "gates", "table_indices", "select_on_true", "name=None"], + "call": ["norm_out=norm_out", "add_input=add_input", "bias_input=bias_input", "gates=gates", "table_indices=table_indices", "select_on_true=select_on_true", "name=name"], + }, + { + "op": "MusaPRelu", + "raw": "musa_p_relu", + "api": "prelu", + "source": "musa_ext/kernels/fusion/musa_prelu_fusion_op.cc", + "args": ["x", "alpha", "name=None"], + "call": ["x=x", "alpha=alpha", "name=name"], + }, + { + "op": "MusaReshapeMatMul", + "raw": "musa_reshape_mat_mul", + "api": "reshape_mat_mul", + "source": "musa_ext/kernels/fusion/musa_reshape_matmul_op.cc", + "args": ["x", "w", "transpose_b=False", "name=None"], + "call": ["x=x", "w=w", "transpose_b=transpose_b", "name=name"], + }, + { + "op": "MusaResourceApplyAdamMixed", + "raw": "musa_resource_apply_adam_mixed", + "api": "resource_apply_adam_mixed", + "source": "musa_ext/kernels/training/musa_applyadam_mixed_op.cc", + "args": ["var", "m", "v", "beta1_power", "beta2_power", "lr", "beta1", "beta2", "epsilon", "grad", "use_locking=False", "use_nesterov=False", "name=None"], + "call": ["var=var", "m=m", "v=v", "beta1_power=beta1_power", "beta2_power=beta2_power", "lr=lr", "beta1=beta1", "beta2=beta2", "epsilon=epsilon", "grad=grad", "use_locking=use_locking", "use_nesterov=use_nesterov", "name=name"], + }, + { + "op": "MusaResourceSparseApplyAdam", + "raw": "musa_resource_sparse_apply_adam", + "api": "resource_sparse_apply_adam", + "source": "musa_ext/kernels/training/musa_apply_sparse_adam_op.cc", + "args": ["var", "m", "v", "beta1_power", "beta2_power", "lr", "beta1", "beta2", "epsilon", "grad", "indices", "use_locking=False", "name=None"], + "call": ["var=var", "m=m", "v=v", "beta1_power=beta1_power", "beta2_power=beta2_power", "lr=lr", "beta1=beta1", "beta2=beta2", "epsilon=epsilon", "grad=grad", "indices=indices", "use_locking=use_locking", "name=name"], + }, + { + "op": "MusaShiftedAffineMap", + "raw": "musa_shifted_affine_map", + "api": "shifted_affine_map", + "source": "musa_ext/kernels/fusion/musa_shifted_affine_map_op.cc", + "args": ["data_left", "mask", "sliced_var_right", "name=None"], + "call": ["data_left=data_left", "mask=mask", "sliced_var_right=sliced_var_right", "name=name"], + }, + { + "op": "MusaTensorDot", + "raw": "musa_tensor_dot", + "api": "tensor_dot", + "source": "musa_ext/kernels/fusion/musa_tensordot_op.cc", + "args": ["a", "b", "axes_a", "axes_b", "name=None"], + "call": ["a=a", "b=b", "axes_a=axes_a", "axes_b=axes_b", "name=name"], + }, + { + "op": "MusaTensorDotBias", + "raw": "musa_tensor_dot_bias", + "api": "tensor_dot_bias", + "source": "musa_ext/kernels/fusion/musa_tensordot_bias_op.cc", + "args": ["a", "b", "bias", "axes_a", "axes_b", "name=None"], + "call": ["a=a", "b=b", "bias=bias", "axes_a=axes_a", "axes_b=axes_b", "name=name"], + }, + { + "op": "MusaTokenMixer", + "raw": "musa_token_mixer", + "api": "token_mixer", + "source": "musa_ext/kernels/fusion/musa_tokenmixer_op.cc", + "args": ["x", "num_T", "num_H", "d_k", "name=None"], + "call": ["x=x", "num_T=num_T", "num_H=num_H", "d_k=d_k", "name=name"], + }, + { + "op": "ResourceApplyNadam", + "raw": "ResourceApplyNadam", + "api": "resource_apply_nadam", + "source": "musa_ext/kernels/training/musa_resource_apply_nadam_op.cc", + "args": ["var", "m", "v", "beta1_power", "beta2_power", "lr", "beta1", "beta2", "epsilon", "grad", "use_locking=False", "name=None"], + "call": ["var=var", "m=m", "v=v", "beta1_power=beta1_power", "beta2_power=beta2_power", "lr=lr", "beta1=beta1", "beta2=beta2", "epsilon=epsilon", "grad=grad", "use_locking=use_locking", "name=name"], + }, +] + +HEADER = """# Copyright 2026 The TensorFlow MUSA Authors. All Rights Reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n# ==============================================================================\n\n""" + + +def _registered_ops_from_source() -> set[str]: + pattern = re.compile(r'REGISTER_OP\("([^"]+)"\)') + registered = set() + for path in KERNELS_ROOT.rglob("*.cc"): + registered.update(pattern.findall(path.read_text(encoding="utf-8"))) + return registered + + +def _validate_definitions() -> None: + source_ops = _registered_ops_from_source() + defined_ops = {entry["op"] for entry in OP_DEFINITIONS} + missing = sorted(source_ops - defined_ops) + stale = sorted(defined_ops - source_ops) + if missing or stale: + messages = [] + if missing: + messages.append(f"missing from OP_DEFINITIONS: {missing}") + if stale: + messages.append(f"not found in REGISTER_OP sources: {stale}") + raise SystemExit("Custom op manifest is out of sync: " + "; ".join(messages)) + + for entry in OP_DEFINITIONS: + args_src = ", ".join(entry["args"]) + ast.parse(f"def {entry['api']}({args_src}):\n pass\n") + + +def _format_call(raw_name: str, call_args: list[str]) -> str: + if len(call_args) <= 3: + return f" return raw_ops.{raw_name}({', '.join(call_args)})\n" + lines = [f" return raw_ops.{raw_name}(\n"] + lines.extend(f" {arg},\n" for arg in call_args) + lines.append(" )\n") + return "".join(lines) + + +def _generate_ops() -> str: + lines = [HEADER, '"""Generated public wrappers for MUSA extension ops."""\n\n', "from . import raw_ops\n\n\n"] + for entry in OP_DEFINITIONS: + lines.append(f"def {entry['api']}({', '.join(entry['args'])}):\n") + lines.append(_format_call(entry["raw"], entry["call"])) + lines.append("\n\n") + + lines.append("__all__ = [\n") + for api_name in sorted(entry["api"] for entry in OP_DEFINITIONS): + lines.append(f' "{api_name}",\n') + lines.append("]\n") + return "".join(lines) + + +def _generate_manifest() -> str: + serializable = [ + { + "op": entry["op"], + "raw": entry["raw"], + "api": entry["api"], + "source": entry["source"], + } + for entry in OP_DEFINITIONS + ] + lines = [HEADER, '"""Manifest of custom MUSA TensorFlow ops exposed to Python."""\n\n'] + lines.append("CUSTOM_OPS = ") + lines.append(pprint.pformat(serializable, width=88, sort_dicts=False)) + lines.append("\n\n") + lines.append("CUSTOM_OP_NAMES = tuple(entry[\"op\"] for entry in CUSTOM_OPS)\n") + lines.append("RAW_OP_NAMES = tuple(entry[\"raw\"] for entry in CUSTOM_OPS)\n") + lines.append("PUBLIC_API_NAMES = tuple(entry[\"api\"] for entry in CUSTOM_OPS)\n") + return "".join(lines) + + +def main() -> None: + _validate_definitions() + GENERATED_OPS.write_text(_generate_ops(), encoding="utf-8") + OP_MANIFEST.write_text(_generate_manifest(), encoding="utf-8") + + +if __name__ == "__main__": + main()