Fix for issue #22974. Cast intermediate results to float before adding and casting the result to the destination type. Avoids half+half operator ambiguity. (#22994)
This commit is contained in:
@@ -184,13 +184,15 @@ static __global__ void ggml_cuda_ar_kernel(
|
|||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int k = 0; k < ELEMS_PER_VEC; ++k) {
|
for (int k = 0; k < ELEMS_PER_VEC; ++k) {
|
||||||
const T_wire d_low = ggml_cuda_cast<T_wire>(sendbuf[off + k]);
|
const T_wire d_low = ggml_cuda_cast<T_wire>(sendbuf[off + k]);
|
||||||
recvbuf[off + k] = ggml_cuda_cast<T_dst>(d_low) + ggml_cuda_cast<T_dst>(wire[k]);
|
recvbuf[off + k] = ggml_cuda_cast<T_dst>(
|
||||||
|
ggml_cuda_cast<float>(d_low) + ggml_cuda_cast<float>(wire[k]));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if (bid == 0 && tid < count - tail) {
|
if (bid == 0 && tid < count - tail) {
|
||||||
const T_wire d_low = ggml_cuda_cast<T_wire>(sendbuf[tail + tid]);
|
const T_wire d_low = ggml_cuda_cast<T_wire>(sendbuf[tail + tid]);
|
||||||
recvbuf[tail + tid] =
|
recvbuf[tail + tid] = ggml_cuda_cast<T_dst>(
|
||||||
ggml_cuda_cast<T_dst>(d_low) + ggml_cuda_cast<T_dst>(host_other[tail + tid]);
|
ggml_cuda_cast<float>(d_low) +
|
||||||
|
ggml_cuda_cast<float>(host_other[tail + tid]));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -210,7 +212,8 @@ static __global__ void ggml_cuda_ar_add_kernel(
|
|||||||
const int nt = gridDim.x * blockDim.x;
|
const int nt = gridDim.x * blockDim.x;
|
||||||
for (int i = tid; i < count; i += nt) {
|
for (int i = tid; i < count; i += nt) {
|
||||||
const T_src d_low = ggml_cuda_cast<T_src>(dst[i]);
|
const T_src d_low = ggml_cuda_cast<T_src>(dst[i]);
|
||||||
dst[i] = ggml_cuda_cast<T_dst>(d_low) + ggml_cuda_cast<T_dst>(src[i]);
|
dst[i] = ggml_cuda_cast<T_dst>(
|
||||||
|
ggml_cuda_cast<float>(d_low) + ggml_cuda_cast<float>(src[i]));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user