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
1 change: 1 addition & 0 deletions musa_ext/kernels/array/musa_ZerosLike_op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@ REGISTER_MUSA_ZEROS_LIKE(double);
REGISTER_MUSA_ZEROS_LIKE(int32);
REGISTER_MUSA_ZEROS_LIKE(int64);
REGISTER_MUSA_ZEROS_LIKE(bool);
REGISTER_MUSA_ZEROS_LIKE(bfloat16);

#undef REGISTER_MUSA_ZEROS_LIKE

Expand Down
72 changes: 71 additions & 1 deletion musa_ext/kernels/array/musa_diag_part_kernel.mu
Original file line number Diff line number Diff line change
Expand Up @@ -44,5 +44,75 @@ template void MusaDiagPartkernelLauncher<Eigen::half>(musaStream_t, uint64_t,
template void MusaDiagPartkernelLauncher<Eigen::bfloat16>(
musaStream_t, uint64_t, const Eigen::bfloat16*, Eigen::bfloat16*);

// ==========================================
// MatrixDiagPartV3 kernel
// ==========================================
// Handles batched [..., M, N] inputs with diagonal offset k.
// For k >= 0: extracts super-diagonal k; for k < 0: sub-diagonal.
// Supports both scalar k and range k = [k_min, k_max] (multiple diagonals).
template <typename T>
__global__ void MusaMatrixDiagPartV3Kernel(
const T* __restrict__ input, T* __restrict__ output,
const T padding_value,
int64 batch_size, int64 M, int64 N,
int k_min, int k_max,
int64 num_diags, int64 max_diag_len) {
int64 tid = (int64)blockIdx.x * blockDim.x + threadIdx.x;
int64 total = batch_size * num_diags * max_diag_len;
if (tid >= total) return;

int64 b = tid / (num_diags * max_diag_len);
int64 rem = tid % (num_diags * max_diag_len);
int64 d = rem / max_diag_len; // d=0 => k_max diagonal
int64 i = rem % max_diag_len;

int k = k_max - (int)d;
int64 k_abs = (k < 0) ? (int64)(-k) : (int64)k;
int64 diag_len = (M < N ? M : N) - k_abs;

if (diag_len <= 0 || i >= diag_len) {
output[tid] = padding_value;
return;
}

int64 row = i + (k < 0 ? k_abs : 0LL);
int64 col = i + (k > 0 ? (int64)k : 0LL);
output[tid] = input[b * M * N + row * N + col];
}

template <typename T>
void MusaMatrixDiagPartV3KernelLauncher(musaStream_t stream, int64 batch_size,
int64 M, int64 N, int k_min, int k_max,
int64 num_diags, int64 max_diag_len,
const T padding_value, const T* input,
T* output) {
int64 total = batch_size * num_diags * max_diag_len;
if (total == 0) return;
const int block_size = 256;
const int grid_size =
static_cast<int>((total + block_size - 1) / block_size);
MusaMatrixDiagPartV3Kernel<T><<<grid_size, block_size, 0, stream>>>(
input, output, padding_value, batch_size, M, N, k_min, k_max, num_diags,
max_diag_len);
}

template void MusaMatrixDiagPartV3KernelLauncher<float>(musaStream_t, int64,
int64, int64, int, int, int64, int64, const float, const float*, float*);
template void MusaMatrixDiagPartV3KernelLauncher<double>(musaStream_t, int64,
int64, int64, int, int, int64, int64, const double, const double*, double*);
template void MusaMatrixDiagPartV3KernelLauncher<int32>(musaStream_t, int64,
int64, int64, int, int, int64, int64, const int32, const int32*, int32*);
template void MusaMatrixDiagPartV3KernelLauncher<long long>(musaStream_t, int64,
int64, int64, int, int, int64, int64, const long long, const long long*,
long long*);
template void MusaMatrixDiagPartV3KernelLauncher<long>(musaStream_t, int64,
int64, int64, int, int, int64, int64, const long, const long*, long*);
template void MusaMatrixDiagPartV3KernelLauncher<Eigen::half>(musaStream_t,
int64, int64, int64, int, int, int64, int64, const Eigen::half,
const Eigen::half*, Eigen::half*);
template void MusaMatrixDiagPartV3KernelLauncher<Eigen::bfloat16>(musaStream_t,
int64, int64, int64, int, int, int64, int64, const Eigen::bfloat16,
const Eigen::bfloat16*, Eigen::bfloat16*);

} // namespace musa
} // namespace tensorflow
} // namespace tensorflow
97 changes: 97 additions & 0 deletions musa_ext/kernels/array/musa_diag_part_op.cc
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
#include "../utils_op.h"
#include "tensorflow/core/framework/bfloat16.h"
#include "tensorflow/core/framework/tensor_shape.h"

namespace tensorflow {
namespace musa {
Expand All @@ -8,6 +9,13 @@ template <typename T>
void MusaDiagPartkernelLauncher(musaStream_t stream, uint64_t size, const T* in,
T* out);

template <typename T>
void MusaMatrixDiagPartV3KernelLauncher(musaStream_t stream, int64 batch_size,
int64 M, int64 N, int k_min, int k_max,
int64 num_diags, int64 max_diag_len,
const T padding_value, const T* input,
T* output);

template <typename T>
class MusaDiagPartOp : public MusaOpKernel {
/*
Expand Down Expand Up @@ -65,5 +73,94 @@ REGISTER_MUSA_DIAG_PART(int64);
REGISTER_MUSA_DIAG_PART(Eigen::half);
REGISTER_MUSA_DIAG_PART(bfloat16);

// ==========================================
// MatrixDiagPartV3 Op
// ==========================================
// Handles tf.linalg.diag_part(x) which maps to MatrixDiagPartV3.
// Supports batched [..., M, N] inputs with scalar or range k.
template <typename T>
class MusaMatrixDiagPartV3Op : public MusaOpKernel {
public:
explicit MusaMatrixDiagPartV3Op(OpKernelConstruction* ctx)
: MusaOpKernel(ctx) {}

void Compute(OpKernelContext* ctx) override {
const Tensor& input = ctx->input(0);
const Tensor& k_tensor = ctx->input(1);
const Tensor& padding_tensor = ctx->input(2);

// Parse k (scalar or 1-D length-2 vector)
int k_min, k_max;
const int32* k_data = k_tensor.flat<int32>().data();
if (k_tensor.NumElements() == 1) {
k_min = k_max = static_cast<int>(k_data[0]);
} else {
k_min = static_cast<int>(k_data[0]);
k_max = static_cast<int>(k_data[1]);
}
OP_REQUIRES(ctx, k_min <= k_max,
errors::InvalidArgument("k[0] must be <= k[1], got k=[", k_min,
",", k_max, "]"));

const T padding_value = padding_tensor.scalar<T>()();
const bool scalar_diag = (k_min == k_max);
const int num_diags = k_max - k_min + 1;

const TensorShape& in_shape = input.shape();
const int ndims = in_shape.dims();
OP_REQUIRES(
ctx, ndims >= 2,
errors::InvalidArgument("Input must be at least 2D, got rank ", ndims));

const int64 M = in_shape.dim_size(ndims - 2);
const int64 N = in_shape.dim_size(ndims - 1);

// Compute max_diag_len across all requested diagonals
int64 max_diag_len = 0;
for (int k = k_min; k <= k_max; ++k) {
int64 dl = std::min(M, N) - static_cast<int64>(std::abs(k));
if (dl > max_diag_len) max_diag_len = dl;
}
OP_REQUIRES(ctx, max_diag_len > 0,
errors::InvalidArgument("k is out of bounds for matrix [", M,
", ", N, "]"));

// Build output shape: [...batch_dims..., (num_diags if range,)
// max_diag_len]
TensorShape out_shape;
for (int i = 0; i < ndims - 2; ++i) {
out_shape.AddDim(in_shape.dim_size(i));
}
if (!scalar_diag) out_shape.AddDim(static_cast<int64>(num_diags));
out_shape.AddDim(max_diag_len);

Tensor* output = nullptr;
OP_REQUIRES_OK(ctx, ctx->allocate_output(0, out_shape, &output));

const int64 batch_size = input.NumElements() / (M * N);

MUSA_OP_REQUIRES_MUDNN_HANDLE(ctx);
auto& handle = GetHandleByCtx(ctx);
musaStream_t stream = reinterpret_cast<musaStream_t>(handle.GetStream());

MusaMatrixDiagPartV3KernelLauncher<T>(
stream, batch_size, M, N, k_min, k_max, static_cast<int64>(num_diags),
max_diag_len, padding_value, input.flat<T>().data(),
output->flat<T>().data());
}
};

#define REGISTER_MUSA_MATRIX_DIAG_PART_V3(TYPE) \
REGISTER_KERNEL_BUILDER( \
Name("MatrixDiagPartV3").Device("MUSA").TypeConstraint<TYPE>("T"), \
MusaMatrixDiagPartV3Op<TYPE>)

REGISTER_MUSA_MATRIX_DIAG_PART_V3(float);
REGISTER_MUSA_MATRIX_DIAG_PART_V3(double);
REGISTER_MUSA_MATRIX_DIAG_PART_V3(int32);
REGISTER_MUSA_MATRIX_DIAG_PART_V3(int64);
REGISTER_MUSA_MATRIX_DIAG_PART_V3(Eigen::half);
REGISTER_MUSA_MATRIX_DIAG_PART_V3(bfloat16);

} // namespace musa
} // namespace tensorflow
2 changes: 2 additions & 0 deletions musa_ext/kernels/array/musa_pack_op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -396,6 +396,7 @@ void MusaUnpackOp<bfloat16>::LaunchUnpackSingleForType(
// Register Pack operators
REGISTER_MUSA_PACK_KERNELS(float)
REGISTER_MUSA_PACK_KERNELS(double)
REGISTER_MUSA_PACK_KERNELS(int32)
REGISTER_MUSA_PACK_KERNELS(int64)
REGISTER_MUSA_PACK_KERNELS(Eigen::half)
REGISTER_MUSA_PACK_KERNELS(bfloat16)
Expand All @@ -405,6 +406,7 @@ REGISTER_MUSA_PACK_KERNELS(uint8)
// Register Unpack operators
REGISTER_MUSA_UNPACK_KERNELS(float)
REGISTER_MUSA_UNPACK_KERNELS(double)
REGISTER_MUSA_UNPACK_KERNELS(int32)
REGISTER_MUSA_UNPACK_KERNELS(int64)
REGISTER_MUSA_UNPACK_KERNELS(Eigen::half)
REGISTER_MUSA_UNPACK_KERNELS(bfloat16)
Expand Down
51 changes: 48 additions & 3 deletions musa_ext/kernels/array/musa_pad_op.cc
100755 → 100644
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,11 @@ void nd_pad_kernel_launcher_uint8(const uint8_t *, uint8_t *, const int,
const int64_t *, const int64_t *,
const uint8_t, const int64_t,
const musaStream_t);
void nd_pad_kernel_launcher_uint16(const uint16_t *, uint16_t *, const int,
const int64_t *, const int64_t *,
const int64_t *, const int64_t *,
const uint16_t, const int64_t,
const musaStream_t);
}

namespace {
Expand All @@ -48,6 +53,8 @@ struct DoubleTag {};
struct Int32Tag {};
struct Int64Tag {};
struct Uint8Tag {};
struct HalfTag {};
struct BFloat16Tag {};
struct UnsupportedTag {};

template <typename T>
Expand Down Expand Up @@ -75,6 +82,14 @@ template <>
struct TypeTag<uint8_t> {
using type = Uint8Tag;
};
template <>
struct TypeTag<Eigen::half> {
using type = HalfTag;
};
template <>
struct TypeTag<bfloat16> {
using type = BFloat16Tag;
};

#define DEFINE_PAD_LAUNCHER_IMPL(T, TAG, SUFFIX) \
void CallPadLauncherImpl( \
Expand All @@ -94,6 +109,37 @@ DEFINE_PAD_LAUNCHER_IMPL(int64_t, Int64Tag, int64)
DEFINE_PAD_LAUNCHER_IMPL(uint8_t, Uint8Tag, uint8)

#undef DEFINE_PAD_LAUNCHER_IMPL

// half and bfloat16 dispatched via uint16 kernel (zero = 0x0000 for both)
void CallPadLauncherImpl(const Eigen::half *input_data,
Eigen::half *output_data, const int dims,
const int64_t *in_dims, const int64_t *out_dims,
const int64_t *pad_before, const int64_t *pad_after,
const Eigen::half pad_value,
const int64_t total_out_elements,
const musaStream_t stream, HalfTag) {
uint16_t pad_bits;
memcpy(&pad_bits, &pad_value, sizeof(uint16_t));
nd_pad_kernel_launcher_uint16(reinterpret_cast<const uint16_t *>(input_data),
reinterpret_cast<uint16_t *>(output_data), dims,
in_dims, out_dims, pad_before, pad_after,
pad_bits, total_out_elements, stream);
}

void CallPadLauncherImpl(const bfloat16 *input_data, bfloat16 *output_data,
const int dims, const int64_t *in_dims,
const int64_t *out_dims, const int64_t *pad_before,
const int64_t *pad_after, const bfloat16 pad_value,
const int64_t total_out_elements,
const musaStream_t stream, BFloat16Tag) {
uint16_t pad_bits;
memcpy(&pad_bits, &pad_value, sizeof(uint16_t));
nd_pad_kernel_launcher_uint16(reinterpret_cast<const uint16_t *>(input_data),
reinterpret_cast<uint16_t *>(output_data), dims,
in_dims, out_dims, pad_before, pad_after,
pad_bits, total_out_elements, stream);
}

void CallPadLauncherImpl(const void *, void *, const int, const int64_t *,
const int64_t *, const int64_t *, const int64_t *,
const int64_t, const int64_t, const musaStream_t,
Expand All @@ -107,9 +153,6 @@ void CallPadLauncher(const T *input_data, T *output_data, const int dims,
const int64_t *pad_before, const int64_t *pad_after,
const T pad_value, const int64_t total_out_elements,
const musaStream_t stream) {
static_assert(!std::is_same<typename TypeTag<T>::type, UnsupportedTag>::value,
"Unsupported type for nd_pad_kernel_launcher");

CallPadLauncherImpl(input_data, output_data, dims, in_dims, out_dims,
pad_before, pad_after, pad_value, total_out_elements,
stream, typename TypeTag<T>::type());
Expand Down Expand Up @@ -246,6 +289,8 @@ REGISTER_MUSA_PAD_TYPE(int32);
REGISTER_MUSA_PAD_TYPE(int64);
REGISTER_MUSA_PAD_TYPE(double);
REGISTER_MUSA_PAD_TYPE(uint8);
REGISTER_MUSA_PAD_TYPE(Eigen::half);
REGISTER_MUSA_PAD_TYPE(bfloat16);

#undef REGISTER_MUSA_PAD_TYPE
} // namespace musa
Expand Down
4 changes: 4 additions & 0 deletions musa_ext/kernels/array/musa_pad_op.mu
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ INSTANTIATE_ND_PAD_KERNEL(double)
INSTANTIATE_ND_PAD_KERNEL(int32_t)
INSTANTIATE_ND_PAD_KERNEL(int64_t)
INSTANTIATE_ND_PAD_KERNEL(uint8_t)
INSTANTIATE_ND_PAD_KERNEL(uint16_t)

#undef INSTANTIATE_ND_PAD_KERNEL

Expand All @@ -87,5 +88,8 @@ DEFINE_ND_PAD_LAUNCHER(int32_t, int32)
DEFINE_ND_PAD_LAUNCHER(int64_t, int64)
DEFINE_ND_PAD_LAUNCHER(uint8_t, uint8)

// uint16_t launcher used by both half and bfloat16 (same bit-width, zero = 0x0000)
DEFINE_ND_PAD_LAUNCHER(uint16_t, uint16)

#undef DEFINE_ND_PAD_LAUNCHER
} // extern "C"
1 change: 1 addition & 0 deletions musa_ext/kernels/array/musa_tensorlist_fromtensor_op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,7 @@ REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(bfloat16);
REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(int32);
REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(int64);
REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(uint8);
REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(bool);

#undef REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR

Expand Down
1 change: 1 addition & 0 deletions musa_ext/kernels/array/musa_where_kernel.mu
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,7 @@ INSTANTIATE_SELECT_FLAGGED_ALL(int16);
INSTANTIATE_SELECT_FLAGGED_ALL(uint16);
INSTANTIATE_SELECT_FLAGGED_ALL(int32);
INSTANTIATE_SELECT_FLAGGED_ALL(int64);
INSTANTIATE_SELECT_FLAGGED_ALL(Eigen::half);
INSTANTIATE_SELECT_FLAGGED_ALL(bfloat16);
#undef INSTANTIATE_SELECT_FLAGGED
#undef INSTANTIATE_SELECT_FLAGGED_ALL
Expand Down
1 change: 1 addition & 0 deletions musa_ext/kernels/array/musa_where_op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -115,6 +115,7 @@ REGISTER_MUSA_WHERE_OP(uint16);
REGISTER_MUSA_WHERE_OP(int32);
REGISTER_MUSA_WHERE_OP(int64);
REGISTER_MUSA_WHERE_OP(bfloat16);
REGISTER_MUSA_WHERE_OP(Eigen::half);
REGISTER_MUSA_WHERE_OP(bool);

#undef REGISTER_MUSA_WHERE_OP
Expand Down
6 changes: 3 additions & 3 deletions musa_ext/kernels/math/musa_cast_op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ namespace tensorflow {
namespace musa {

extern "C" void LaunchFloatToBFloat16Copy(const float* src, void* dst,
int64_t n, musaStream_t stream);
int64_t n, musaStream_t stream);

class MusaCastOp : public MusaOpKernel {
public:
Expand Down Expand Up @@ -44,8 +44,7 @@ class MusaCastOp : public MusaOpKernel {
return;
}

if (external_src_dtype_ == DT_FLOAT &&
external_dst_dtype_ == DT_BFLOAT16) {
if (external_src_dtype_ == DT_FLOAT && external_dst_dtype_ == DT_BFLOAT16) {
LaunchFloatToBFloat16Copy(
inp.flat<float>().data(),
reinterpret_cast<void*>(output->flat<bfloat16>().data()),
Expand Down Expand Up @@ -108,6 +107,7 @@ REGISTER_CAST_MUSA(int32, int32);
REGISTER_CAST_MUSA(int32, int64);
REGISTER_CAST_MUSA(int32, Eigen::half);
REGISTER_CAST_MUSA(int32, bfloat16);
REGISTER_CAST_MUSA(int32, float);
REGISTER_CAST_MUSA(int32, double);

REGISTER_CAST_MUSA(int64, bool);
Expand Down
2 changes: 2 additions & 0 deletions musa_ext/kernels/math/musa_log_op.cc
100755 → 100644
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
#include <mudnn.h>

#include <type_traits>

#include "../utils_op.h"
#include "tensorflow/core/framework/op_kernel.h"
#include "tensorflow/core/framework/register_types.h"
Expand Down
5 changes: 5 additions & 0 deletions musa_ext/kernels/math/musa_select_op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -356,6 +356,11 @@ REGISTER_SELECT(int64);
REGISTER_SELECT(bool);
REGISTER_SELECT(Eigen::half);
REGISTER_SELECT(Eigen::bfloat16);
REGISTER_SELECT(int8);
REGISTER_SELECT(int16);
REGISTER_SELECT(uint8);
REGISTER_SELECT(uint16);
REGISTER_SELECT(uint32);

#undef REGISTER_SELECT

Expand Down
Loading
Loading