Skip to content

Commit d00edfc

Browse files
Albert/disable soft placement config for test (#286)
* Fix(DisableSoftPlacementConfig)
1 parent 3ef12ac commit d00edfc

29 files changed

Lines changed: 755 additions & 114 deletions

musa_ext/kernels/array/musa_ZerosLike_op.cc

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@ REGISTER_MUSA_ZEROS_LIKE(double);
4444
REGISTER_MUSA_ZEROS_LIKE(int32);
4545
REGISTER_MUSA_ZEROS_LIKE(int64);
4646
REGISTER_MUSA_ZEROS_LIKE(bool);
47+
REGISTER_MUSA_ZEROS_LIKE(bfloat16);
4748

4849
#undef REGISTER_MUSA_ZEROS_LIKE
4950

musa_ext/kernels/array/musa_diag_part_kernel.mu

Lines changed: 71 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -44,5 +44,75 @@ template void MusaDiagPartkernelLauncher<Eigen::half>(musaStream_t, uint64_t,
4444
template void MusaDiagPartkernelLauncher<Eigen::bfloat16>(
4545
musaStream_t, uint64_t, const Eigen::bfloat16*, Eigen::bfloat16*);
4646

47+
// ==========================================
48+
// MatrixDiagPartV3 kernel
49+
// ==========================================
50+
// Handles batched [..., M, N] inputs with diagonal offset k.
51+
// For k >= 0: extracts super-diagonal k; for k < 0: sub-diagonal.
52+
// Supports both scalar k and range k = [k_min, k_max] (multiple diagonals).
53+
template <typename T>
54+
__global__ void MusaMatrixDiagPartV3Kernel(
55+
const T* __restrict__ input, T* __restrict__ output,
56+
const T padding_value,
57+
int64 batch_size, int64 M, int64 N,
58+
int k_min, int k_max,
59+
int64 num_diags, int64 max_diag_len) {
60+
int64 tid = (int64)blockIdx.x * blockDim.x + threadIdx.x;
61+
int64 total = batch_size * num_diags * max_diag_len;
62+
if (tid >= total) return;
63+
64+
int64 b = tid / (num_diags * max_diag_len);
65+
int64 rem = tid % (num_diags * max_diag_len);
66+
int64 d = rem / max_diag_len; // d=0 => k_max diagonal
67+
int64 i = rem % max_diag_len;
68+
69+
int k = k_max - (int)d;
70+
int64 k_abs = (k < 0) ? (int64)(-k) : (int64)k;
71+
int64 diag_len = (M < N ? M : N) - k_abs;
72+
73+
if (diag_len <= 0 || i >= diag_len) {
74+
output[tid] = padding_value;
75+
return;
76+
}
77+
78+
int64 row = i + (k < 0 ? k_abs : 0LL);
79+
int64 col = i + (k > 0 ? (int64)k : 0LL);
80+
output[tid] = input[b * M * N + row * N + col];
81+
}
82+
83+
template <typename T>
84+
void MusaMatrixDiagPartV3KernelLauncher(musaStream_t stream, int64 batch_size,
85+
int64 M, int64 N, int k_min, int k_max,
86+
int64 num_diags, int64 max_diag_len,
87+
const T padding_value, const T* input,
88+
T* output) {
89+
int64 total = batch_size * num_diags * max_diag_len;
90+
if (total == 0) return;
91+
const int block_size = 256;
92+
const int grid_size =
93+
static_cast<int>((total + block_size - 1) / block_size);
94+
MusaMatrixDiagPartV3Kernel<T><<<grid_size, block_size, 0, stream>>>(
95+
input, output, padding_value, batch_size, M, N, k_min, k_max, num_diags,
96+
max_diag_len);
97+
}
98+
99+
template void MusaMatrixDiagPartV3KernelLauncher<float>(musaStream_t, int64,
100+
int64, int64, int, int, int64, int64, const float, const float*, float*);
101+
template void MusaMatrixDiagPartV3KernelLauncher<double>(musaStream_t, int64,
102+
int64, int64, int, int, int64, int64, const double, const double*, double*);
103+
template void MusaMatrixDiagPartV3KernelLauncher<int32>(musaStream_t, int64,
104+
int64, int64, int, int, int64, int64, const int32, const int32*, int32*);
105+
template void MusaMatrixDiagPartV3KernelLauncher<long long>(musaStream_t, int64,
106+
int64, int64, int, int, int64, int64, const long long, const long long*,
107+
long long*);
108+
template void MusaMatrixDiagPartV3KernelLauncher<long>(musaStream_t, int64,
109+
int64, int64, int, int, int64, int64, const long, const long*, long*);
110+
template void MusaMatrixDiagPartV3KernelLauncher<Eigen::half>(musaStream_t,
111+
int64, int64, int64, int, int, int64, int64, const Eigen::half,
112+
const Eigen::half*, Eigen::half*);
113+
template void MusaMatrixDiagPartV3KernelLauncher<Eigen::bfloat16>(musaStream_t,
114+
int64, int64, int64, int, int, int64, int64, const Eigen::bfloat16,
115+
const Eigen::bfloat16*, Eigen::bfloat16*);
116+
47117
} // namespace musa
48-
} // namespace tensorflow
118+
} // namespace tensorflow

musa_ext/kernels/array/musa_diag_part_op.cc

Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
#include "../utils_op.h"
22
#include "tensorflow/core/framework/bfloat16.h"
3+
#include "tensorflow/core/framework/tensor_shape.h"
34

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

12+
template <typename T>
13+
void MusaMatrixDiagPartV3KernelLauncher(musaStream_t stream, int64 batch_size,
14+
int64 M, int64 N, int k_min, int k_max,
15+
int64 num_diags, int64 max_diag_len,
16+
const T padding_value, const T* input,
17+
T* output);
18+
1119
template <typename T>
1220
class MusaDiagPartOp : public MusaOpKernel {
1321
/*
@@ -65,5 +73,94 @@ REGISTER_MUSA_DIAG_PART(int64);
6573
REGISTER_MUSA_DIAG_PART(Eigen::half);
6674
REGISTER_MUSA_DIAG_PART(bfloat16);
6775

76+
// ==========================================
77+
// MatrixDiagPartV3 Op
78+
// ==========================================
79+
// Handles tf.linalg.diag_part(x) which maps to MatrixDiagPartV3.
80+
// Supports batched [..., M, N] inputs with scalar or range k.
81+
template <typename T>
82+
class MusaMatrixDiagPartV3Op : public MusaOpKernel {
83+
public:
84+
explicit MusaMatrixDiagPartV3Op(OpKernelConstruction* ctx)
85+
: MusaOpKernel(ctx) {}
86+
87+
void Compute(OpKernelContext* ctx) override {
88+
const Tensor& input = ctx->input(0);
89+
const Tensor& k_tensor = ctx->input(1);
90+
const Tensor& padding_tensor = ctx->input(2);
91+
92+
// Parse k (scalar or 1-D length-2 vector)
93+
int k_min, k_max;
94+
const int32* k_data = k_tensor.flat<int32>().data();
95+
if (k_tensor.NumElements() == 1) {
96+
k_min = k_max = static_cast<int>(k_data[0]);
97+
} else {
98+
k_min = static_cast<int>(k_data[0]);
99+
k_max = static_cast<int>(k_data[1]);
100+
}
101+
OP_REQUIRES(ctx, k_min <= k_max,
102+
errors::InvalidArgument("k[0] must be <= k[1], got k=[", k_min,
103+
",", k_max, "]"));
104+
105+
const T padding_value = padding_tensor.scalar<T>()();
106+
const bool scalar_diag = (k_min == k_max);
107+
const int num_diags = k_max - k_min + 1;
108+
109+
const TensorShape& in_shape = input.shape();
110+
const int ndims = in_shape.dims();
111+
OP_REQUIRES(
112+
ctx, ndims >= 2,
113+
errors::InvalidArgument("Input must be at least 2D, got rank ", ndims));
114+
115+
const int64 M = in_shape.dim_size(ndims - 2);
116+
const int64 N = in_shape.dim_size(ndims - 1);
117+
118+
// Compute max_diag_len across all requested diagonals
119+
int64 max_diag_len = 0;
120+
for (int k = k_min; k <= k_max; ++k) {
121+
int64 dl = std::min(M, N) - static_cast<int64>(std::abs(k));
122+
if (dl > max_diag_len) max_diag_len = dl;
123+
}
124+
OP_REQUIRES(ctx, max_diag_len > 0,
125+
errors::InvalidArgument("k is out of bounds for matrix [", M,
126+
", ", N, "]"));
127+
128+
// Build output shape: [...batch_dims..., (num_diags if range,)
129+
// max_diag_len]
130+
TensorShape out_shape;
131+
for (int i = 0; i < ndims - 2; ++i) {
132+
out_shape.AddDim(in_shape.dim_size(i));
133+
}
134+
if (!scalar_diag) out_shape.AddDim(static_cast<int64>(num_diags));
135+
out_shape.AddDim(max_diag_len);
136+
137+
Tensor* output = nullptr;
138+
OP_REQUIRES_OK(ctx, ctx->allocate_output(0, out_shape, &output));
139+
140+
const int64 batch_size = input.NumElements() / (M * N);
141+
142+
MUSA_OP_REQUIRES_MUDNN_HANDLE(ctx);
143+
auto& handle = GetHandleByCtx(ctx);
144+
musaStream_t stream = reinterpret_cast<musaStream_t>(handle.GetStream());
145+
146+
MusaMatrixDiagPartV3KernelLauncher<T>(
147+
stream, batch_size, M, N, k_min, k_max, static_cast<int64>(num_diags),
148+
max_diag_len, padding_value, input.flat<T>().data(),
149+
output->flat<T>().data());
150+
}
151+
};
152+
153+
#define REGISTER_MUSA_MATRIX_DIAG_PART_V3(TYPE) \
154+
REGISTER_KERNEL_BUILDER( \
155+
Name("MatrixDiagPartV3").Device("MUSA").TypeConstraint<TYPE>("T"), \
156+
MusaMatrixDiagPartV3Op<TYPE>)
157+
158+
REGISTER_MUSA_MATRIX_DIAG_PART_V3(float);
159+
REGISTER_MUSA_MATRIX_DIAG_PART_V3(double);
160+
REGISTER_MUSA_MATRIX_DIAG_PART_V3(int32);
161+
REGISTER_MUSA_MATRIX_DIAG_PART_V3(int64);
162+
REGISTER_MUSA_MATRIX_DIAG_PART_V3(Eigen::half);
163+
REGISTER_MUSA_MATRIX_DIAG_PART_V3(bfloat16);
164+
68165
} // namespace musa
69166
} // namespace tensorflow

musa_ext/kernels/array/musa_pack_op.cc

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -396,6 +396,7 @@ void MusaUnpackOp<bfloat16>::LaunchUnpackSingleForType(
396396
// Register Pack operators
397397
REGISTER_MUSA_PACK_KERNELS(float)
398398
REGISTER_MUSA_PACK_KERNELS(double)
399+
REGISTER_MUSA_PACK_KERNELS(int32)
399400
REGISTER_MUSA_PACK_KERNELS(int64)
400401
REGISTER_MUSA_PACK_KERNELS(Eigen::half)
401402
REGISTER_MUSA_PACK_KERNELS(bfloat16)
@@ -405,6 +406,7 @@ REGISTER_MUSA_PACK_KERNELS(uint8)
405406
// Register Unpack operators
406407
REGISTER_MUSA_UNPACK_KERNELS(float)
407408
REGISTER_MUSA_UNPACK_KERNELS(double)
409+
REGISTER_MUSA_UNPACK_KERNELS(int32)
408410
REGISTER_MUSA_UNPACK_KERNELS(int64)
409411
REGISTER_MUSA_UNPACK_KERNELS(Eigen::half)
410412
REGISTER_MUSA_UNPACK_KERNELS(bfloat16)

musa_ext/kernels/array/musa_pad_op.cc

100755100644
Lines changed: 48 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,11 @@ void nd_pad_kernel_launcher_uint8(const uint8_t *, uint8_t *, const int,
4040
const int64_t *, const int64_t *,
4141
const uint8_t, const int64_t,
4242
const musaStream_t);
43+
void nd_pad_kernel_launcher_uint16(const uint16_t *, uint16_t *, const int,
44+
const int64_t *, const int64_t *,
45+
const int64_t *, const int64_t *,
46+
const uint16_t, const int64_t,
47+
const musaStream_t);
4348
}
4449

4550
namespace {
@@ -48,6 +53,8 @@ struct DoubleTag {};
4853
struct Int32Tag {};
4954
struct Int64Tag {};
5055
struct Uint8Tag {};
56+
struct HalfTag {};
57+
struct BFloat16Tag {};
5158
struct UnsupportedTag {};
5259

5360
template <typename T>
@@ -75,6 +82,14 @@ template <>
7582
struct TypeTag<uint8_t> {
7683
using type = Uint8Tag;
7784
};
85+
template <>
86+
struct TypeTag<Eigen::half> {
87+
using type = HalfTag;
88+
};
89+
template <>
90+
struct TypeTag<bfloat16> {
91+
using type = BFloat16Tag;
92+
};
7893

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

96111
#undef DEFINE_PAD_LAUNCHER_IMPL
112+
113+
// half and bfloat16 dispatched via uint16 kernel (zero = 0x0000 for both)
114+
void CallPadLauncherImpl(const Eigen::half *input_data,
115+
Eigen::half *output_data, const int dims,
116+
const int64_t *in_dims, const int64_t *out_dims,
117+
const int64_t *pad_before, const int64_t *pad_after,
118+
const Eigen::half pad_value,
119+
const int64_t total_out_elements,
120+
const musaStream_t stream, HalfTag) {
121+
uint16_t pad_bits;
122+
memcpy(&pad_bits, &pad_value, sizeof(uint16_t));
123+
nd_pad_kernel_launcher_uint16(reinterpret_cast<const uint16_t *>(input_data),
124+
reinterpret_cast<uint16_t *>(output_data), dims,
125+
in_dims, out_dims, pad_before, pad_after,
126+
pad_bits, total_out_elements, stream);
127+
}
128+
129+
void CallPadLauncherImpl(const bfloat16 *input_data, bfloat16 *output_data,
130+
const int dims, const int64_t *in_dims,
131+
const int64_t *out_dims, const int64_t *pad_before,
132+
const int64_t *pad_after, const bfloat16 pad_value,
133+
const int64_t total_out_elements,
134+
const musaStream_t stream, BFloat16Tag) {
135+
uint16_t pad_bits;
136+
memcpy(&pad_bits, &pad_value, sizeof(uint16_t));
137+
nd_pad_kernel_launcher_uint16(reinterpret_cast<const uint16_t *>(input_data),
138+
reinterpret_cast<uint16_t *>(output_data), dims,
139+
in_dims, out_dims, pad_before, pad_after,
140+
pad_bits, total_out_elements, stream);
141+
}
142+
97143
void CallPadLauncherImpl(const void *, void *, const int, const int64_t *,
98144
const int64_t *, const int64_t *, const int64_t *,
99145
const int64_t, const int64_t, const musaStream_t,
@@ -107,9 +153,6 @@ void CallPadLauncher(const T *input_data, T *output_data, const int dims,
107153
const int64_t *pad_before, const int64_t *pad_after,
108154
const T pad_value, const int64_t total_out_elements,
109155
const musaStream_t stream) {
110-
static_assert(!std::is_same<typename TypeTag<T>::type, UnsupportedTag>::value,
111-
"Unsupported type for nd_pad_kernel_launcher");
112-
113156
CallPadLauncherImpl(input_data, output_data, dims, in_dims, out_dims,
114157
pad_before, pad_after, pad_value, total_out_elements,
115158
stream, typename TypeTag<T>::type());
@@ -246,6 +289,8 @@ REGISTER_MUSA_PAD_TYPE(int32);
246289
REGISTER_MUSA_PAD_TYPE(int64);
247290
REGISTER_MUSA_PAD_TYPE(double);
248291
REGISTER_MUSA_PAD_TYPE(uint8);
292+
REGISTER_MUSA_PAD_TYPE(Eigen::half);
293+
REGISTER_MUSA_PAD_TYPE(bfloat16);
249294

250295
#undef REGISTER_MUSA_PAD_TYPE
251296
} // namespace musa

musa_ext/kernels/array/musa_pad_op.mu

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,7 @@ INSTANTIATE_ND_PAD_KERNEL(double)
6464
INSTANTIATE_ND_PAD_KERNEL(int32_t)
6565
INSTANTIATE_ND_PAD_KERNEL(int64_t)
6666
INSTANTIATE_ND_PAD_KERNEL(uint8_t)
67+
INSTANTIATE_ND_PAD_KERNEL(uint16_t)
6768

6869
#undef INSTANTIATE_ND_PAD_KERNEL
6970

@@ -87,5 +88,8 @@ DEFINE_ND_PAD_LAUNCHER(int32_t, int32)
8788
DEFINE_ND_PAD_LAUNCHER(int64_t, int64)
8889
DEFINE_ND_PAD_LAUNCHER(uint8_t, uint8)
8990

91+
// uint16_t launcher used by both half and bfloat16 (same bit-width, zero = 0x0000)
92+
DEFINE_ND_PAD_LAUNCHER(uint16_t, uint16)
93+
9094
#undef DEFINE_ND_PAD_LAUNCHER
9195
} // extern "C"

musa_ext/kernels/array/musa_tensorlist_fromtensor_op.cc

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -134,6 +134,7 @@ REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(bfloat16);
134134
REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(int32);
135135
REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(int64);
136136
REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(uint8);
137+
REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR(bool);
137138

138139
#undef REGISTER_MUSA_TENSOR_LIST_FROM_TENSOR
139140

musa_ext/kernels/array/musa_where_kernel.mu

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -129,6 +129,7 @@ INSTANTIATE_SELECT_FLAGGED_ALL(int16);
129129
INSTANTIATE_SELECT_FLAGGED_ALL(uint16);
130130
INSTANTIATE_SELECT_FLAGGED_ALL(int32);
131131
INSTANTIATE_SELECT_FLAGGED_ALL(int64);
132+
INSTANTIATE_SELECT_FLAGGED_ALL(Eigen::half);
132133
INSTANTIATE_SELECT_FLAGGED_ALL(bfloat16);
133134
#undef INSTANTIATE_SELECT_FLAGGED
134135
#undef INSTANTIATE_SELECT_FLAGGED_ALL

musa_ext/kernels/array/musa_where_op.cc

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -115,6 +115,7 @@ REGISTER_MUSA_WHERE_OP(uint16);
115115
REGISTER_MUSA_WHERE_OP(int32);
116116
REGISTER_MUSA_WHERE_OP(int64);
117117
REGISTER_MUSA_WHERE_OP(bfloat16);
118+
REGISTER_MUSA_WHERE_OP(Eigen::half);
118119
REGISTER_MUSA_WHERE_OP(bool);
119120

120121
#undef REGISTER_MUSA_WHERE_OP

musa_ext/kernels/math/musa_cast_op.cc

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@ namespace tensorflow {
88
namespace musa {
99

1010
extern "C" void LaunchFloatToBFloat16Copy(const float* src, void* dst,
11-
int64_t n, musaStream_t stream);
11+
int64_t n, musaStream_t stream);
1212

1313
class MusaCastOp : public MusaOpKernel {
1414
public:
@@ -44,8 +44,7 @@ class MusaCastOp : public MusaOpKernel {
4444
return;
4545
}
4646

47-
if (external_src_dtype_ == DT_FLOAT &&
48-
external_dst_dtype_ == DT_BFLOAT16) {
47+
if (external_src_dtype_ == DT_FLOAT && external_dst_dtype_ == DT_BFLOAT16) {
4948
LaunchFloatToBFloat16Copy(
5049
inp.flat<float>().data(),
5150
reinterpret_cast<void*>(output->flat<bfloat16>().data()),
@@ -108,6 +107,7 @@ REGISTER_CAST_MUSA(int32, int32);
108107
REGISTER_CAST_MUSA(int32, int64);
109108
REGISTER_CAST_MUSA(int32, Eigen::half);
110109
REGISTER_CAST_MUSA(int32, bfloat16);
110+
REGISTER_CAST_MUSA(int32, float);
111111
REGISTER_CAST_MUSA(int32, double);
112112

113113
REGISTER_CAST_MUSA(int64, bool);

0 commit comments

Comments
 (0)