This guide explains how to write custom MUSA C++ operators that override torch_musa's default ATen implementations.
torchada allows you to override ATen operators at the C++ level for the PrivateUse1 (MUSA) dispatch key. This is useful when you need:
- Better performance than the default torch_musa implementation
- Custom behavior for specific operators
- Workarounds for torch_musa bugs
C++ extensions are automatically loaded on MUSA platform when torchada is imported.
Edit src/torchada/csrc/musa_ops.mu:
#include "ops.h"
#include <ATen/musa/MUSAContext.h>
namespace torchada {
template <typename scalar_t>
__global__ void my_kernel(scalar_t* output, const scalar_t* input, int64_t n) {
int64_t idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx < n) {
output[idx] = /* your computation */;
}
}
at::Tensor my_op_impl(const at::Tensor& self) {
log_op_call("my_op");
auto input = self.contiguous();
auto output = at::empty_like(input);
if (input.numel() == 0) return output;
musaStream_t stream = at::musa::getCurrentMUSAStream();
const int64_t n = input.numel();
const int threads = 256;
const int blocks = (n + threads - 1) / threads;
AT_DISPATCH_FLOATING_TYPES(input.scalar_type(), "my_op", [&] {
my_kernel<scalar_t><<<blocks, threads, 0, stream>>>(
output.data_ptr<scalar_t>(),
input.data_ptr<scalar_t>(),
n);
});
return output;
}
} // namespace torchada
TORCH_LIBRARY_IMPL(aten, PrivateUse1, m) {
// Check env var at registration time - allows disabling via
// TORCHADA_DISABLE_OP_OVERRIDE_my_op=1
if (torchada::is_override_enabled("my_op")) {
m.impl("my_op", torchada::my_op_impl);
}
}TORCHADA_DEBUG_CPP_OPS=1 python -c "
import torch
import torchada
x = torch.randn(1000, device='cuda')
y = torch.neg(x) # Should print '[torchada] neg called'
print('Result:', y.cpu()[:5])
"| File | Purpose |
|---|---|
src/torchada/csrc/ops.h |
Header with utilities (log_op_call, is_override_enabled) |
src/torchada/csrc/ops.cpp |
Python bindings, C++-only implementations, and CUDA-compatible APIs |
src/torchada/csrc/musa_ops.mu |
MUSA kernel implementations |
torchada's C++ extension provides CUDA-compatible implementations of some torch_musa memory management APIs:
These functions are automatically injected into torch.cuda.memory and allow CUDA code using memory pools to work transparently on MUSA:
_cuda_beginAllocateCurrentThreadToPool(device, mempool_id)- Begin allocating memory from current thread to a memory pool_cuda_endAllocateToPool(device, mempool_id)- End allocating memory to a memory pool_cuda_releasePool(device, mempool_id)- Release a memory pool
Usage example:
import torchada
import torch
# This works transparently on MUSA - no code changes needed
from torch.cuda.memory import _cuda_beginAllocateCurrentThreadToPool
from torch.cuda.memory import _cuda_endAllocateToPool
from torch.cuda.memory import _cuda_releasePool
# Use the functions as in CUDA code
device = 0
pool_id = torch.cuda.graph_pool_handle()
_cuda_beginAllocateCurrentThreadToPool(device, pool_id)
# ... allocations ...
_cuda_endAllocateToPool(device, pool_id)
_cuda_releasePool(device, pool_id)| Variable | Description |
|---|---|
TORCHADA_CPP_OPS_VERBOSE=1 |
Show compilation output |
TORCHADA_DEBUG_CPP_OPS=1 |
Log operator calls to stdout |
TORCHADA_DISABLE_OP_OVERRIDE_<NAME>=1 |
Disable specific operator override |
MTGPU_TARGET=mp_XX |
Override GPU architecture detection |
To disable a specific operator override at runtime, set the environment variable before importing torchada:
# Disable the 'neg' operator override, use torch_musa's default instead
TORCHADA_DISABLE_OP_OVERRIDE_neg=1 python my_script.pyImportant: The operator name in the environment variable should match the name passed to is_override_enabled() in the C++ code. For example, if the code uses is_override_enabled("neg"), set TORCHADA_DISABLE_OP_OVERRIDE_neg=1.
This check happens at registration time (when the extension is loaded), not at runtime. Once the extension is loaded, the operator registrations are fixed.
torchada auto-detects the GPU architecture using musaInfo:
| GPU | Compute Capability | Architecture |
|---|---|---|
| MTT S80 | 2.1 | mp_21 |
| MTT S4000 | 2.2 | mp_22 |
| MTT S5000 | 3.1 | mp_31 |
Override with: export MTGPU_TARGET=mp_22
When overriding an operator, don't call the same operator:
// BAD - causes infinite recursion
at::Tensor bad_neg_impl(const at::Tensor& self) {
return -self; // Calls aten::neg again!
}
// GOOD - use lower-level primitives
at::Tensor good_neg_impl(const at::Tensor& self) {
auto output = at::empty_like(self);
// Launch custom kernel or use in-place ops
return output;
}at::Tensor my_impl(const at::Tensor& self) {
auto input = self.contiguous(); // Ensure contiguous
if (input.numel() == 0) {
return at::empty_like(input); // Handle empty tensors
}
// ... kernel launch
}musaError_t err = musaGetLastError();
if (err != musaSuccess) {
TORCH_CHECK(false, "MUSA kernel failed: ", musaGetErrorString(err));
}AT_DISPATCH_ALL_TYPES_AND2(
at::ScalarType::Half, at::ScalarType::BFloat16,
input.scalar_type(), "my_kernel", [&] {
my_kernel<scalar_t><<<blocks, threads, 0, stream>>>(...);
});Any ATen operator with a PrivateUse1 dispatch can be overridden. Common categories:
abs, neg, exp, exp2, log, log2, log10, sqrt, rsqrt, ceil, floor, round, trunc, sign, sin, cos, tan, asin, acos, atan, sinh, cosh, tanh, sigmoid, erf, erfc, reciprocal, bitwise_not
add, sub, mul, div, pow, fmod, remainder, maximum, minimum, atan2, bitwise_and, bitwise_or, bitwise_xor, logical_and, logical_or, logical_xor
sum, prod, mean, std, var, max, min, argmax, argmin, all, any, norm, logsumexp
mm, bmm, addmm, addmv, addr, matmul, dot, mv, ger, linear
relu, relu_, leaky_relu, gelu, silu, mish, hardswish, hardsigmoid, softplus, softshrink, threshold
batch_norm, layer_norm, group_norm, instance_norm, local_response_norm
max_pool1d, max_pool2d, max_pool3d, avg_pool1d, avg_pool2d, avg_pool3d, adaptive_max_pool2d, adaptive_avg_pool2d
conv1d, conv2d, conv3d, conv_transpose1d, conv_transpose2d, conv_transpose3d
copy_, clone, contiguous, fill_, zero_, ones_like, zeros_like, empty_like
index, index_put_, gather, scatter, scatter_add, masked_fill, masked_select, where
view, reshape, transpose, permute, squeeze, unsqueeze, expand, repeat, cat, stack, split, chunk
To find the exact operator signature, use:
import torch
# Search for specific operator:
for s in torch._C._jit_get_all_schemas():
if 'neg' in str(s):
print(s)See src/torchada/csrc/musa_ops.mu for a complete example of an aten::neg override (commented out by default — uncomment the m.impl() line to activate).
- Verify kernel is called: Set
TORCHADA_DEBUG_CPP_OPS=1 - Check compilation: Set
TORCHADA_CPP_OPS_VERBOSE=1 - Clear cache:
rm -rf ~/.cache/torch_extensions/*/torchada_cpp_ops - Check architecture: Run
musaInfo | grep "compute capability"