You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何编写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;
}

关键说明

  1. 四维索引映射逻辑
    不用强行把四维对应到三维的block/grid维度,而是先计算每个线程的全局线性ID,再通过除法和取余反推出四维坐标(w,i,j,k),这种方式适配所有维度扩展场景,逻辑更清晰。

  2. 内存容量限制
    N=220时,四维张量的元素数是(220)4=280,单精度浮点下内存需求约68GB,普通消费级GPU根本装不下,实际开发必须用更小的N,或者拆分张量分批处理。

  3. 网格维度限制
    CUDA对网格各维度有最大值限制:x轴最大为2^31-1,y、z轴最大为65535。代码中计算网格大小时已经考虑了这个限制,避免超出范围导致内核启动失败。

  4. 替代实现思路
    如果你想更直观地拆分维度,可以把第四个维度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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.26 16:33:09