如何理解 mma_tile 函数

NPU基本架构(含buffer/MTE相关)

基本架构详细说明参考基本架构-昇腾社区,最重要的是其中的硬件架构图

Buffer速查表

Buffer描述
L1 BufferL1缓冲区,可暂存Cube计算单元需要反复使用的一些数据从而减少从总线读写的次数。
L0A/L0B BufferCube指令的输入。
L0C BufferCube指令的输出,但进行累加计算的时候,也是输入的一部分。
BT BufferBiasTable Buffer,存放矩阵计算中的Bias。

地址修饰符速查表

修饰符含义
__gm__全局内存 (HBM/DDR)
__cbuf__L1 缓存 (Cube Buffer)
__ca__L0A 缓存 (CUBE A矩阵专用)
__cb__L0B 缓存 (CUBE B矩阵专用)
__cc__L0C 缓存 (CUBE 累加器)

搬运单元速查表 (GM = global memory)

搬运单元搬运通路
MTE1L1 -> L0A/L0B/BT
MTE2GM -> L1/L0A/L0B
FixPipeL0C -> GM/L1

流水线枚举速查表

枚举描述
PIPE_S标量计算流水线
PIPE_V向量计算流水线
PIPE_M内存操作流水线
PIPE_MTE1内存传输引擎1流水线
PIPE_MTE2内存传输引擎2流水线
PIPE_MTE3内存传输引擎3流水线
PIPE_ALL所有流水线
PIPE_FIX固定功能流水线

同步控制

参考NPU架构版本300x-同步控制。主要需要理解的是,AI Core内部的执行单元(如MTE2搬运单元、Cube计算单元等)是异步并行的,因此需要同步控制来避免数据覆盖等冲突。set_flagwait_flag就是为了实现这一点。

个人对set_flagwait_flag的理解是:wait_flag会等待set_flag发出“你可以继续了”的信号,类似于消息传递;或者说,wait_flag是加一把锁,等待set_flag来解锁。信号由第一个参数发给第二个参数。对于set_flag,指令发给第一个参数;对于wait_flag,指令发给第二个参数。

典型的场景有:前一次计算未完成时,不允许搬运数据到计算输入缓存;前一次计算的结果未完全搬出输出缓存前,不允许进行新的计算。这个同步控制流程图以时间轴的形式展现了上述例子。

从另一个角度看,搬入/计算/搬出单元有独立的指令队列,可以把set_flagwait_flag视为发送给它们的指令,其中wait_flag会被阻塞直到对应的set_flag执行。这个指令队列示意图非常直观生动。值得一提的是,从这张图可以隐约看出一个DAG关系,可体现指令间的依赖和执行顺序。

根据CANN文档,set_flagwait_flag的使用有以下注意事项,虽然不一定完全适用于IR项目,但可供参考:

  • 建议通过AllocEventID或者FetchEventID接口获取EventID,以确保其合法性和有效性。
  • EventID的数量有限,使用后应立即调用ReleaseEventID释放资源,避免EventID耗尽,影响系统正常运行。
  • SetFlagWaitFlag必须成对使用,且SetFlagWaitFlag的参数必须完全一致,才表示同一个EventID。

此外,个人认为,函数间的事件传递也很重要,需要谨慎设计和处理。

INTRINSIC宏是什么

mma_tile函数中,有很多类似INTRINSIC(wait_flag, PIPE_M, PIPE_MTE1, event)的调用,但INTRINSIC宏其实并不复杂。它的定义是:

// bishengir/lib/Template/include/CPUDebug/CPUDebugUtils.h:160
#ifdef ENABLE_CPU_TRACE_INTRINSIC
    #define INTRINSIC(NAME, ...) printIntrinsic(TO_STRING(NAME), __VA_ARGS__)
#else
    #define INTRINSIC(NAME, ...) NAME(__VA_ARGS__)
#endif

在NPU上,INTRINSIC会直接调用对应的硬件指令。例如,INTRINSIC(wait_flag, PIPE_M, PIPE_MTE1, event)会转化为wait_flag(PIPE_M, PIPE_MTE1, event),其中wait_flag是硬件提供的。而在非NPU上,INTRINSIC只打印这个调用。可以简单理解为,这个宏的第一个参数是一个函数,而它会调用这个函数,并把后续所有参数传入。

mma_tile分析

以下是mma_tile函数的分析。所有中文注释均为个人理解,不保证正确。

// bishengir/lib/Template/lib/Cube/LocalMmad.cpp:314
template <typename SRC_TYPE, // 输入(矩阵a和b)的数据类型
          typename DST_TYPE, // 输入(矩阵c)的数据类型
          typename BIAS_TYPE, // bias的数据类型
          bool TA, // 是否在计算前转置矩阵a
          bool TB, // 是否在计算前转置矩阵b
          bool HF32, // 是否使用HF32模式,无需理解
          bool I4> // 是否为int4数据类型
__aicore__ __attribute__((always_inline)) void
mma_tile(memref_t<__cbuf__ SRC_TYPE, 4> *ma, // 这里的__cbuf__表示使用L1缓存的地址空间,详见本页开头的地址修饰符速查表
         memref_t<__cbuf__ SRC_TYPE, 4> *mb,
         bool init, // 是否在计算前将L0C归零
         int64_t m, int64_t k, int64_t n, // 矩阵维度
         memref_t<__cbuf__ BIAS_TYPE, 4> *bias, // 可选的bias
         memref_t<__cc__ DST_TYPE, 4> *mc, // 同上,__cc__表示使用L0C缓存的地址空间
         // 以下是各种事件等,将在需要时说明
         int64_t mmad_l1_wait_l1a_event, // 用于MTE2告知MTE1矩阵a搬运完成
         int64_t mmad_l1_wait_l1b_event, // 用于MTE2告知MTE1矩阵b搬运完成
         int64_t l1a_wait_mmad_l1_event, // 期望MTE1告知MTE2矩阵a使用完毕
         int64_t l1b_wait_mmad_l1_event, // 期望MTE1告知MTE2矩阵b使用完毕
         int64_t kloop_db_cond,
         int64_t back_pipe_m_pipe_mte1_db_event0, // 用于以及期望Cube告知MTE1上半区数据使用完成
         int64_t back_pipe_m_pipe_mte1_db_event1, // 用于以及期望Cube告知MTE1下半区数据使用完成
         uint8_t unit_flag
) {
  // 如果矩阵是空的,则可以无需计算直接返回,但需要处理外部传入的信号
  // 由搬运单元速查表知,MTE2是GM->L1,MTE1是L1->L0
  if (m == 0 || k == 0 || n == 0) {
    if (mmad_l1_wait_l1a_event != -1) {
      // 等待MTE2告知MTE1:从GM搬运矩阵a到L1已完成
      INTRINSIC(wait_flag, PIPE_MTE2, PIPE_MTE1, mmad_l1_wait_l1a_event);
    }
    if (mmad_l1_wait_l1b_event != -1) {
      // 等待MTE2告知MTE1:从GM搬运矩阵b到L1已完成
      INTRINSIC(wait_flag, PIPE_MTE2, PIPE_MTE1, mmad_l1_wait_l1b_event);
    }
    if (l1a_wait_mmad_l1_event != -1) {
      // MTE1告知MTE2:L1中的矩阵a已使用完毕,可以覆盖
      INTRINSIC(set_flag, PIPE_MTE1, PIPE_MTE2, l1a_wait_mmad_l1_event);
    }
    if (l1b_wait_mmad_l1_event != -1) {
      // MTE1告知MTE2:L1中的矩阵b已使用完毕,可以覆盖
      INTRINSIC(set_flag, PIPE_MTE1, PIPE_MTE2, l1b_wait_mmad_l1_event);
    }
    return;
  }

  // int4类型不支持转置矩阵a
  if constexpr (I4) {
    static_assert(TA == 0 && "i4 mmad doesn't support transpose A");
  }

  // 这里的类型可以读作:指向【存储在L0C中的、DST_TYPE类型的数据】的指针
  __cc__ DST_TYPE *mc_ptr = mc->aligned + mc->offset;
  __ca__ SRC_TYPE *l0a_base = reinterpret_cast<__ca__ SRC_TYPE *>((uintptr_t)0);
  __cb__ SRC_TYPE *l0b_base = reinterpret_cast<__cb__ SRC_TYPE *>((uintptr_t)0);

  // 以下很长一段内容都在计算矩阵大小和数据分块,因为L0缓存太小,通常无法把数据一次性从L1搬运到L0
  const int64_t k_actual = k;
  const int64_t m0 = TA ? (ma->sizes[0] * ma->sizes[3]) : (ma->sizes[1] * ma->sizes[2]);
  const int64_t n0 = TB ? (mb->sizes[1] * mb->sizes[2]) : (mb->sizes[0] * mb->sizes[3]);
  const int64_t mn_max = I4 ? (m0 > 2 * n0 ? m0 : 2 * n0) : (m0 > n0 ? m0 : n0);
  const int64_t l0c_m_size = mc->sizes[1] * mc->sizes[2];

  bool k_direction_align = (sizeof(SRC_TYPE) == 4) && TA;
  int64_t k_ceil = CEIL_FACTOR(k, L1_ALIGN_BYTES / sizeof(SRC_TYPE));
  k_ceil = k_direction_align ? CEIL_FACTOR(k_ceil, FRACTAL_BLOCK_NUM) : k_ceil;
  const int64_t k_ceil_b = I4 ? CEIL_FACTOR(2 * k, L1_ALIGN_BYTES / sizeof(SRC_TYPE)) : 0;
  const int64_t elem_num_per_block = L1_ALIGN_BYTES / sizeof(SRC_TYPE);
  const int64_t max_align_value = elem_num_per_block > 16 ? elem_num_per_block : 16;

  int64_t n_ceil = n;
  if constexpr ((sizeof(SRC_TYPE) == 1)  && !I4) {
    n_ceil = TB ? CEIL_FACTOR(n_ceil, FRACTAL_BLOCK_NUM)
                : CEIL_FACTOR(n_ceil, elem_num_per_block);
  }
  n_ceil = k_direction_align ? CEIL_FACTOR(n, 16) : n_ceil;

  // 尝试使用 ping-pong 策略,即把L0缓存分成两半,其中一半作为计算输入时可以同时把数据往另一半搬运
  // compute k_part, namely how much k can be put into L0, it is main block
  // aligned tile size
  bool enable_double_buffer = true;
  int64_t l0ab_pingpong_buffer_len =
      L0AB_BUFFER_BYTES / 2 / sizeof(SRC_TYPE);
  int64_t k_part = FLOOR_FACTOR(l0ab_pingpong_buffer_len /
                                    CEIL_FACTOR(mn_max, max_align_value),
                                max_align_value);
  // L0的一半放不下一个分块,无法使用 ping-pong 策略
  if (k_part == 0) {
    enable_double_buffer = false;
    l0ab_pingpong_buffer_len = L0AB_BUFFER_BYTES / sizeof(SRC_TYPE);
    k_part = FLOOR_FACTOR(l0ab_pingpong_buffer_len /
      CEIL_FACTOR(mn_max, max_align_value),
      max_align_value);
  }
  int64_t k_part_b = (I4 && (m0 < n0)) ? (2 * k_part) : k_part;

  if (k_part == 0) {
    trap(); // 硬件异常
  }

  // 接下来,Cube(M)需要在ping-pong分区中的数据使用完毕时告知MTE1,因此需要两个分别对应两个半区的事件
  // 如果外部传入了这两个event,则使用对应的event,否则使用默认的EVENT_ID0和EVENT_ID1
  // 如果使用默认的EVENT_ID,则需要在循环前进行初始化,以便在第一次循环前告知MTE1可以搬运数据
  if (back_pipe_m_pipe_mte1_db_event0 == -1) {
    INTRINSIC(set_flag, PIPE_M, PIPE_MTE1, EVENT_ID0);
  }
  if (back_pipe_m_pipe_mte1_db_event1 == -1) {
    INTRINSIC(set_flag, PIPE_M, PIPE_MTE1, EVENT_ID1);
  }

  // 对不同分块进行循环
  const int64_t k_part_loop = (k_actual + k_part - 1) / k_part;
  for (int64_t k_part_idx = 0; k_part_idx < k_part_loop; k_part_idx++) {
    const bool is_outer_k_start = k_part_idx == 0;
    const bool is_outer_k_end = k_part_idx + 1 == k_part_loop;
    const bool init_c = init && is_outer_k_start;

    // compute k_part_actual and k_part_ceil
    // k_part_ceil means the aligned k tile size in L0 for k_part_idx iteration
    // k_part_actual means that unaligned k tile size in L0 for k_part_idx
    // iteration
    int64_t k_part_ceil =
        !is_outer_k_end ? k_part : (k_ceil - k_part_idx * k_part);
    int64_t k_part_actual =
        !is_outer_k_end ? k_part : (k_actual - k_part_idx * k_part);

    int64_t k_part_ceil_b = 0;
    int64_t k_part_actual_b = 0;
    if constexpr (I4) {
      k_part_ceil_b =
          !is_outer_k_end ? k_part_b : (k_ceil_b - k_part_idx * k_part_b);
      k_part_actual_b =
          !is_outer_k_end ? k_part_b : (2 * k_actual - k_part_idx * k_part_b);
    }

    // 如果不使用double buffer,则直接使用0作为ping_pong_id并获取对应事件
    /* TODO: 但这里是不是有个bug:如果调用者期望通过back_pipe_m_pipe_mte1_db_event1来阻止下半分区的数据覆盖,而mma_tile恰好没有使用double buffer,那么这个event就会被直接忽略,导致数据被提前覆盖 */
    int64_t ping_pong_id = 0;
    if (enable_double_buffer) {
      ping_pong_id = (kloop_db_cond != -1)
                         ? (kloop_db_cond * k_part_loop + k_part_idx) % 2
                         : k_part_idx % 2;
    }
    int64_t m_mte1_event_id = (ping_pong_id == 0)
                                  ? (back_pipe_m_pipe_mte1_db_event0 != -1
                                         ? back_pipe_m_pipe_mte1_db_event0
                                         : EVENT_ID0)
                                  : (back_pipe_m_pipe_mte1_db_event1 != -1
                                         ? back_pipe_m_pipe_mte1_db_event1
                                         : EVENT_ID1);
    __ca__ SRC_TYPE *l0a_buf =
        l0a_base + ping_pong_id * l0ab_pingpong_buffer_len;
    __cb__ SRC_TYPE *l0b_buf =
        l0b_base + ping_pong_id * l0ab_pingpong_buffer_len;

    // 等待Cube告知MTE1:此半区中的数据已使用完毕,可以覆盖
    INTRINSIC(wait_flag, PIPE_M, PIPE_MTE1, m_mte1_event_id);

    // load matrix A from L1 to L0A
    // 如果这是第一轮循环,则等待MTE2告知MTE1:从GM搬运矩阵a到L1已完成
    if (is_outer_k_start && mmad_l1_wait_l1a_event != -1) {
      INTRINSIC(wait_flag, PIPE_MTE2, PIPE_MTE1, mmad_l1_wait_l1a_event);
    }

    // 发送搬运数据指令
    load_l1_to_l0a<SRC_TYPE, TA, I4>(l0a_buf, ma, k_part_idx, k_part, k_part_ceil,
                                     k_part_loop, k_part_actual, l0c_m_size);

    // 如果这是最后一轮循环,则MTE1告知MTE2:L1中的矩阵a已使用完毕,可以覆盖
    if (is_outer_k_end && l1a_wait_mmad_l1_event != -1) {
      INTRINSIC(set_flag, PIPE_MTE1, PIPE_MTE2, l1a_wait_mmad_l1_event);
    }

    // load matrix B from L1 to L0B
    // 如果这是第一轮循环,则等待MTE2告知MTE1:从GM搬运矩阵b到L1已完成
    if (is_outer_k_start && mmad_l1_wait_l1b_event != -1) {
      INTRINSIC(wait_flag, PIPE_MTE2, PIPE_MTE1, mmad_l1_wait_l1b_event);
    }

    // int4类型的特殊处理
    if constexpr (I4) {
      load_l1_to_l0b<SRC_TYPE, TB, I4>(l0b_buf, mb, k_part_idx, k_part_b, k_part_ceil_b,
                                       k_part_loop, k_part_actual_b, n);
    } else {
      load_l1_to_l0b<SRC_TYPE, TB, I4>(l0b_buf, mb, k_part_idx, k_part, k_part_ceil,
                                       k_part_loop, k_part_actual, n);
    }

    // 如果这是最后一轮循环,则MTE1告知MTE2:L1中的矩阵b已使用完毕,可以覆盖
    if (is_outer_k_end && l1b_wait_mmad_l1_event != -1) {
      INTRINSIC(set_flag, PIPE_MTE1, PIPE_MTE2, l1b_wait_mmad_l1_event);
    }

    // load matrix bias from L1 to BT
    // 值得一提的是,这里没有显式控制Bias数据从L2到L1的同步,原因可能是已经被矩阵a和b的同步保护覆盖了,所以不再单独控制
    if (bias) {
      const uint64_t btAddr = 0;
      __cbuf__ BIAS_TYPE *bias_ptr = bias->aligned + bias->offset;
      int64_t bias_burst_len = CEIL_DIV(n0 * sizeof(BIAS_TYPE), 64);
      int16_t conv_control = (sizeof(BIAS_TYPE) == 2 && sizeof(DST_TYPE) == 4);
      INTRINSIC(copy_cbuf_to_bt, btAddr, bias_ptr,
                conv_control,   // convControl
                1,              // nBurst
                bias_burst_len, // lenBurst
                0,              // sourceGap
                0);             // dstGap
    }

    // mmad
    if constexpr (HF32) {
      set_hf32_ctrl();
    }

    // pipe_m wait for pipe_mte1 instructions to finish
    // 等待MTE1告知Cube:已完成从L1到L0&Bias的搬运
    // 这里个人理解是,MTE1会在前序指令(也就是load_l1_to_l0指令)完成之后再处理后续的set_flag
    // 所以这两句的意义是告诉MTE1在搬运完成后发送信息,让Cube等待这个信息
    // 也引申出一个重要理解:set_flag和load_l1_to_l0这些只是发送指令,不会阻塞到对应操作完成,阻塞需要由flag来显式控制
    INTRINSIC(set_flag, PIPE_MTE1, PIPE_M, m_mte1_event_id);
    INTRINSIC(wait_flag, PIPE_MTE1, PIPE_M, m_mte1_event_id);

    uint8_t unit_flag_mode =
        unit_flag ? (is_outer_k_end ? unit_flag : (uint8_t)0b10) : (uint8_t)0;

    // 计算指令
    if constexpr (I4) {
      if (bias) {
        bool isSourceFromBT = init_c;
        mad_intrin_core_s4(mmad_intrin_args<void, int32_t>{
                           mc_ptr, l0a_buf, l0b_buf, static_cast<uint16_t>(l0c_m_size),
                           static_cast<uint16_t>(2 * k_part_actual), static_cast<uint16_t>(2 * n_ceil),
                           unit_flag_mode, 0, isSourceFromBT, 0});
      } else {
        mad_intrin_core_s4(mmad_intrin_args<void, int32_t>{
                           mc_ptr, l0a_buf, l0b_buf, static_cast<uint16_t>(l0c_m_size),
                           static_cast<uint16_t>(2 * k_part_actual), static_cast<uint16_t>(2 * n_ceil), unit_flag_mode,
                           0, 0, init_c});
      }
    } else {
        if (bias) {
        bool isSourceFromBT = init_c;
        mad_intrin_core(mmad_intrin_args<SRC_TYPE, DST_TYPE>{
            mc_ptr, l0a_buf, l0b_buf, static_cast<uint16_t>(l0c_m_size),
            static_cast<uint16_t>(k_part_actual), static_cast<uint16_t>(n_ceil),
            unit_flag_mode, k_direction_align, /*cmatrixSource=*/isSourceFromBT,
            /*cmatrixInitVal=*/0});
      } else {
        mad_intrin_core(mmad_intrin_args<SRC_TYPE, DST_TYPE>{
            mc_ptr, l0a_buf, l0b_buf, static_cast<uint16_t>(l0c_m_size),
            static_cast<uint16_t>(k_part_actual), static_cast<uint16_t>(n_ceil),
            unit_flag_mode, k_direction_align, /*cmatrixSource=*/0,
            /*cmatrixInitVal=*/init_c});
      }
    }

    // The same l0c address is used. When M/16 * N/16>=10,
    // it is not necessary to insert BAR. M is applicable
    // to V300 and C200 versions.
    if (l0c_m_size / FRACTAL_BLOCK_NUM * n / FRACTAL_BLOCK_NUM < 10) {
      INTRINSIC(pipe_barrier, PIPE_M);
    }

    // pipe_m set next iteration pipe_mte1
    // Cube告知MTE1:此半区数据已使用完毕,可以覆盖
    INTRINSIC(set_flag, PIPE_M, PIPE_MTE1, m_mte1_event_id);

    if constexpr (HF32) {
      set_hf32_ctrl_none();
    }
  }
  // 如果外部没有传入这两个事件,则无法通知外部等待,因此要告知MTE1在进行下一步前等待Cube计算完成
  // p.s. 外部接下来调用的wait_flag不会造成竞争,因为前序指令已经给到MTE1了,外部的指令只会加在最后面
  if (back_pipe_m_pipe_mte1_db_event0 == -1) {
    INTRINSIC(wait_flag, PIPE_M, PIPE_MTE1, EVENT_ID0);
  }
  if (back_pipe_m_pipe_mte1_db_event1 == -1) {
    INTRINSIC(wait_flag, PIPE_M, PIPE_MTE1, EVENT_ID1);
  }
}