如何编写CUDA代码实现N×N×N×N张量相加?
四维张量(N×N×N×N)的CUDA相加实现方案
你已经有了三维张量相加的CUDA代码,现在要扩展到四维,但CUDA的线程块和网格最多只支持三维,核心解决思路是用三维的线程/网格索引组合,映射覆盖四维的张量元素,下面是具体实现和关键说明:
完整实现代码
#include <stdio.h> #include <math.h> #define N (1 << 10) // 注意:N=2^20时四维张量内存需求约68GB,远超普通设备容量,测试用小值 #define BLOCK_SIZE 16 __global__ void add_4d(float* a, float* b, float* c, int n) { // 计算全局线程ID:把三维的block和thread索引转换成全局线性ID long long global_thread_id = (long long)blockIdx.z * gridDim.x * gridDim.y * blockDim.x * blockDim.y * blockDim.z + (long long)blockIdx.y * gridDim.x * blockDim.x * blockDim.y * blockDim.z + (long long)blockIdx.x * blockDim.x * blockDim.y * blockDim.z + (long long)threadIdx.z * blockDim.x * blockDim.y + (long long)threadIdx.y * blockDim.x + threadIdx.x; // 总元素数是n^4,超出范围直接返回 long long total_elements = (long long)n * n * n * n; if (global_thread_id >= total_elements) return; // 从全局ID反推四维坐标(w, i, j, k) long long w = global_thread_id / (n * n * n); long long remaining = global_thread_id % (n * n * n); int i = remaining / (n * n); remaining %= n * n; int j = remaining / n; int k = remaining % n; // 计算线性索引,内存布局为((w*n + i)*n + j)*n + k long long index = ((w * n + i) * n + j) * n + k; c[index] = a[index] + b[index]; } int main() { float *h_a, *h_b, *h_c; // 主机内存指针 float *d_a, *d_b, *d_c; // 设备内存指针 long long size = (long long)N * N * N * N * sizeof(float); // 主机内存分配与初始化 h_a = (float*)malloc(size); h_b = (float*)malloc(size); h_c = (float*)malloc(size); if (!h_a || !h_b || !h_c) { fprintf(stderr, "主机内存分配失败\n"); return 1; } for (long long w = 0; w < N; w++) { for (int i = 0; i < N; i++) { for (int j = 0; j < N; j++) { for (int k = 0; k < N; k++) { long long index = ((w * N + i) * N + j) * N + k; h_a[index] = w + i + j + k; h_b[index] = 4*N - w - i - j - k; } } } } // 设备内存分配 if (cudaMalloc(&d_a, size) != cudaSuccess || cudaMalloc(&d_b, size) != cudaSuccess || cudaMalloc(&d_c, size) != cudaSuccess) { fprintf(stderr, "设备内存分配失败\n"); free(h_a); free(h_b); free(h_c); return 1; } // 主机到设备数据拷贝 cudaMemcpy(d_a, h_a, size, cudaMemcpyHostToDevice); cudaMemcpy(d_b, h_b, size, cudaMemcpyHostToDevice); // 配置线程块和网格大小 dim3 dimBlock(BLOCK_SIZE, BLOCK_SIZE, BLOCK_SIZE); long long threads_per_block = (long long)dimBlock.x * dimBlock.y * dimBlock.z; long long total_threads = (long long)N * N * N * N; // 计算三维网格的各维度大小,优先填满z轴,再y轴,最后x轴(适配CUDA网格维度限制) long long grid_z = total_threads / (threads_per_block * 65535 * 65535); if (total_threads % (threads_per_block * 65535 * 65535) != 0) grid_z++; long long remaining_threads = total_threads / (threads_per_block * grid_z); long long grid_y = remaining_threads / 65535; if (remaining_threads % 65535 != 0) grid_y++; long long grid_x = remaining_threads / grid_y; if (remaining_threads % grid_y != 0) grid_x++; // CUDA网格的y、z轴最大为65535,x轴最大为2^31-1,此处已做限制 dim3 dimGrid((int)grid_x, (int)grid_y, (int)grid_z); // 启动内核并同步(方便排查错误) add_4d<<<dimGrid, dimBlock>>>(d_a, d_b, d_c, N); cudaDeviceSynchronize(); // 设备到主机结果拷贝 cudaMemcpy(h_c, d_c, size, cudaMemcpyDeviceToHost); // 验证结果(小N时可选) bool valid = true; for (long long idx = 0; idx < total_threads; idx++) { if (fabs(h_c[idx] - (h_a[idx] + h_b[idx])) > 1e-5) { valid = false; fprintf(stderr, "结果错误,索引%lld: 计算值%f,预期值%f\n", idx, h_c[idx], h_a[idx]+h_b[idx]); break; } } if (valid) printf("所有元素计算正确\n"); // 释放内存 free(h_a); free(h_b); free(h_c); cudaFree(d_a); cudaFree(d_b); cudaFree(d_c); return 0; }
关键说明
四维索引映射逻辑
不用强行把四维对应到三维的block/grid维度,而是先计算每个线程的全局线性ID,再通过除法和取余反推出四维坐标(w,i,j,k),这种方式适配所有维度扩展场景,逻辑更清晰。内存容量限制
N=220时,四维张量的元素数是(220)4=280,单精度浮点下内存需求约68GB,普通消费级GPU根本装不下,实际开发必须用更小的N,或者拆分张量分批处理。网格维度限制
CUDA对网格各维度有最大值限制:x轴最大为2^31-1,y、z轴最大为65535。代码中计算网格大小时已经考虑了这个限制,避免超出范围导致内核启动失败。替代实现思路
如果你想更直观地拆分维度,可以把第四个维度w分配到网格的x轴,比如:// 网格x轴大小 = 每个w对应的x轴网格数 * w的总数量 int grid_x_per_w = (N + BLOCK_SIZE - 1) / BLOCK_SIZE; dim3 dimGrid(grid_x_per_w * N, (N+BLOCK_SIZE-1)/BLOCK_SIZE, (N+BLOCK_SIZE-1)/BLOCK_SIZE);然后在内核中计算:
int w = blockIdx.x / grid_x_per_w; int i = blockIdx.x % grid_x_per_w * blockDim.x + threadIdx.x; int j = blockIdx.y * blockDim.y + threadIdx.y; int k = blockIdx.z * blockDim.z + threadIdx.z;这种方式适合维度之间独立性较强的场景,但要注意网格x轴不要超出最大值限制。
内容的提问来源于stack exchange,提问作者user366312
相关产品推荐
相关产品推荐

