v3.9 update (#2213)

Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
Yujia Zhai
2025-04-03 02:10:16 -04:00
committed by GitHub
co-authored by yuzhai
parent 6f4921858b
commit 79fc51f4b8
72 changed files with 19875 additions and 459 deletions
@@ -109,7 +109,8 @@ bool BlockCompareEqual(
Element const *ptr_B,
size_t capacity,
int grid_size = 0,
int block_size = 0) {
int block_size = 0,
cudaStream_t stream = nullptr) {
int equal_flag = 1;
int *device_equal_flag = nullptr;
@@ -146,7 +147,9 @@ bool BlockCompareEqual(
dim3 grid(grid_size, 1, 1);
dim3 block(block_size, 1, 1);
kernel::BlockCompareEqual<Element><<< grid, block >>>(device_equal_flag, ptr_A, ptr_B, capacity);
kernel::BlockCompareEqual<Element><<< grid, block, 0, stream >>>(device_equal_flag, ptr_A, ptr_B, capacity);
cudaStreamSynchronize(stream);
if (cudaMemcpy(
&equal_flag,
@@ -175,7 +178,8 @@ bool BlockCompareRelativelyEqual(
Element epsilon,
Element nonzero_floor,
int grid_size = 0,
int block_size = 0) {
int block_size = 0,
cudaStream_t stream = nullptr) {
int equal_flag = 1;
int *device_equal_flag = nullptr;
@@ -212,7 +216,7 @@ bool BlockCompareRelativelyEqual(
dim3 grid(grid_size, 1, 1);
dim3 block(block_size, 1, 1);
kernel::BlockCompareRelativelyEqual<Element><<< grid, block >>>(
kernel::BlockCompareRelativelyEqual<Element><<< grid, block, 0, stream >>>(
device_equal_flag,
ptr_A,
ptr_B,
@@ -221,6 +225,8 @@ bool BlockCompareRelativelyEqual(
nonzero_floor
);
cudaStreamSynchronize(stream);
if (cudaMemcpy(
&equal_flag,
device_equal_flag,
@@ -232,6 +232,8 @@ ComputeType TensorTransformReduce(
workspace, identity, workspace_size, reduce
);
cudaStreamSynchronize(stream);
if (copy_out) {
cudaError_t result = cudaMemcpy(&identity, workspace, sizeof(identity), cudaMemcpyDeviceToHost);
if (result != cudaSuccess) {
@@ -285,6 +287,8 @@ ComputeType TensorTransformReduce(
workspace, identity, workspace_size, reduce
);
cudaStreamSynchronize(stream);
if (copy_out) {
cudaError_t result = cudaMemcpy(&identity, workspace, sizeof(identity), cudaMemcpyDeviceToHost);
if (result != cudaSuccess) {