高阶矩阵乘API流水问题
收藏回复举报
高阶矩阵乘API流水问题
t('forum.solved') 已解决
发表于2024-11-22 16:09:07
0 查看

各位老师,我们最近在使用高阶矩阵乘API的时候,为了尽可能掩盖数据从global memory传输到buffer的开销,我们在使用高阶API的基础上使用双缓冲区,但是我们的代码profile出来显示没有流水,想请教一下,到底是哪里有问题。以下是我们的核心代码

using matmul_t = Matmul<MatmulType<TPosition::VECOUT, CubeFormat::ND, A_T>,
    MatmulType<TPosition::VECOUT, CubeFormat::ND, B_T, true>,
    MatmulType<TPosition::VECIN, CubeFormat::ND, C_T>>;

__aicore__ inline void Process(TPipe &pipe) {

      for (int i = 0; i < repeat; i++) {
        DataCopy(a1, queryGlobal[queryId * K], K);
        mm_0.SetTensorA(a1, isTransA);
        mm_1.SetTensorA(a1, isTransA);
        SetFlag<HardEvent::M_MTE2>(B_M_MTE2_0);
        SetFlag<HardEvent::M_MTE2>(B_M_MTE2_1);
        for (int j = 0; j < canTiling_num; j += 2) {
            WaitFlag<HardEvent::M_MTE2>(B_M_MTE2_0);
            for (int t = 0; t < N; t++) {
                DataCopy(b1_0[t * K], datasetGlobal[can_id_0 * K], K);
            }
            SetFlag<HardEvent::MTE2_M>(B_MTE2_M_0);
            WaitFlag<HardEvent::M_MTE2>(B_M_MTE2_1);
            for (int t = 0; t < N; t++) {
                DataCopy(b1_1[t * K], datasetGlobal[can_id_1 * K], K);
            }
            SetFlag<HardEvent::MTE2_M>(B_MTE2_M_1);
            WaitFlag<HardEvent::MTE2_M>(B_MTE2_M_0);
            mm_0.SetTensorB(b1_0, isTransB);
            mm_0.IterateAll<false>(c1_0);
            SetFlag<HardEvent::M_MTE2>(B_M_MTE2_0);
            DataCopy(disGlobal[i * can_num + N * j], c1_0, N);
            WaitFlag<HardEvent::MTE2_M>(B_MTE2_M_1);
            mm_1.SetTensorB(b1_1, isTransB);
            mm_1.IterateAll<false>(c1_1);
            SetFlag<HardEvent::M_MTE2>(B_M_MTE2_1);
            DataCopy(disGlobal[i * can_num + N * (j + 1)], c1_1, N);
        }
        WaitFlag<HardEvent::M_MTE2>(B_M_MTE2_0);
        WaitFlag<HardEvent::M_MTE2>(B_M_MTE2_1);
      }
    }

另外,我们的硬件是Atlas 300I 推理卡, cann版本为8.0.RC2.alpha003

本帖最后由 匿名用户2024/11/22 16:34:16 编辑

我要发帖子