}
}
+template<>
+void set_rows_cuda<half, int32_t>(ggml_backend_cuda_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
+ const half * src0_d = (const half *)src0->data;
+ const int32_t * src1_d = (const int32_t *)src1->data;
+
+ GGML_TENSOR_BINARY_OP_LOCALS
+
+ cudaStream_t stream = ctx.stream();
+
+
+ if (dst->type == GGML_TYPE_F16) {
+ set_rows_cuda(
+ src0_d, src1_d, (half*)dst->data,
+ ne00, ne01, ne02, ne03,
+ ne10, ne11, ne12, ne13,
+ nb01, nb02, nb03,
+ nb10, nb11, nb12,
+ nb1, nb2, nb3,
+ stream
+ );
+ } else {
+ GGML_ABORT("unsupported type %s", ggml_type_name(dst->type));
+ }
+}
+
+template<>
+void set_rows_cuda<half, int64_t>(ggml_backend_cuda_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
+ const half * src0_d = (const half *)src0->data;
+ const int64_t * src1_d = (const int64_t *)src1->data;
+
+ GGML_TENSOR_BINARY_OP_LOCALS
+
+ cudaStream_t stream = ctx.stream();
+
+
+ if (dst->type == GGML_TYPE_F16) {
+ set_rows_cuda(
+ src0_d, src1_d, (half*)dst->data,
+ ne00, ne01, ne02, ne03,
+ ne10, ne11, ne12, ne13,
+ nb01, nb02, nb03,
+ nb10, nb11, nb12,
+ nb1, nb2, nb3,
+ stream
+ );
+ } else {
+ GGML_ABORT("unsupported type %s", ggml_type_name(dst->type));
+ }
+}
+
void ggml_cuda_op_set_rows(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
const ggml_tensor * src0 = dst->src[0];
const ggml_tensor * src1 = dst->src[1];
- GGML_ASSERT(src0->type == GGML_TYPE_F32);
+ GGML_ASSERT(src0->type == GGML_TYPE_F32 || (src0->type == GGML_TYPE_F16 && dst->type == GGML_TYPE_F16));
GGML_ASSERT(src1->type == GGML_TYPE_I64 || src1->type == GGML_TYPE_I32);
- if (src1->type == GGML_TYPE_I64) {
- set_rows_cuda<float, int64_t>(ctx, src0, src1, dst);
+ if (src0->type == GGML_TYPE_F32) {
+ if (src1->type == GGML_TYPE_I64) {
+ set_rows_cuda<float, int64_t>(ctx, src0, src1, dst);
+ } else {
+ set_rows_cuda<float, int32_t>(ctx, src0, src1, dst);
+ }
+ } else if (src0->type == GGML_TYPE_F16) {
+ if (src1->type == GGML_TYPE_I64) {
+ set_rows_cuda<half, int64_t>(ctx, src0, src1, dst);
+ } else {
+ set_rows_cuda<half, int32_t>(ctx, src0, src1, dst);
+ }
} else {
- set_rows_cuda<float, int32_t>(ctx, src0, src1, dst);
+ GGML_ABORT("unsupported type %s", ggml_type_name(src0->type));
}
}