调试模块DFX

硬件背景

device_print

device_print是Triton框架在昇腾NPU上提供的一个设备端调试工具,允许开发者在算子内核执行过程中直接打印标量/矢量信息,核心流程如下:

flowchart LR
    subgraph Host[Host 侧流程]
        A[Host Launcher] -->|1. 传递打印缓冲区| B[Kernel 执行]
        B -->|2. Kernel 返回| C[读取缓冲区]
        C -->|3. 解析并打印| D[终端输出]
    end

    subgraph Code[代码实现]
        E[BiSheng 头文件<br/>内置打印逻辑] -->|自动抽取| F[triton-ascend<br/>集成]
        F -->|调用| A
    end

关键硬件资源约束

  • UB打印缓冲区:每个aicore固定分配16 KB空间用于数据暂存,且同一个aicore内的所有打印操作共享这16 KB缓冲区,写满后新数据提示warning,大小超过缓冲区最大值后丢弃。

  • 多核并发:每个aicore独立执行内核代码,最终Host侧呈现每个核的打印结果。

算法原理

实现原理涉及Triton Ascend、AscendNPU IR、毕昇编译器三部分配合,实现该功能主要以AscendNPU IR为重点展开说明。

Triton Ascend

生成初始.ttadapter IR过程中会将triton侧的tl.device_print转换成func.call @triton_print_*接口。

AscendNPU IR

接收到.ttadapter IR之后在AscendNPU IR阶段主要会经历如下变换:

AdaptTritonKernel

func.call @triton_print_*接口转换成hfusion.print接口。

// Before AdaptTritonKernel
%reinterpret_cast = memref.reinterpret_cast %arg2 to offset: [0], sizes: [8], strides: [1] : memref<?xi64> to memref<8xi64, strided<[1]>>
%alloc = memref.alloc() : memref<8xi64>
memref.copy %reinterpret_cast, %alloc : memref<8xi64, strided<[1]>> to memref<8xi64>
%0 = bufferization.to_tensor %alloc restrict writable : memref<8xi64>
call @triton_print_0(%0) : (tensor<8xi64>) -> ()

// After AdaptTritonKernel
%reinterpret_cast = memref.reinterpret_cast %arg2 to offset: [0], sizes: [8], strides: [1] : memref<?xi64> to memref<8xi64, strided<[1]>>
%alloc = memref.alloc() : memref<8xi64>
memref.copy %reinterpret_cast, %alloc : memref<8xi64, strided<[1]>> to memref<8xi64>
%0 = bufferization.to_tensor %alloc restrict writable : memref<8xi64>
hfusion.print " x: " {hex = false} %0 : tensor<8xi64>

HFusionToHIVM

hfusion.print接口转换成hivm.hir.debug接口。

// Before ConvertHFusionToHIVM
%reinterpret_cast = memref.reinterpret_cast %arg3 to offset: [0], sizes: [8], strides: [1] : memref<?xi64> to memref<8xi64, strided<[1]>>
%alloc = memref.alloc() : memref<8xi64>
memref.copy %reinterpret_cast, %alloc : memref<8xi64, strided<[1]>> to memref<8xi64>
%0 = bufferization.to_tensor %alloc restrict writable : memref<8xi64>
hfusion.print " x: " {hex = false} %0 : tensor<8xi64>

// After ConvertHFusionToHIVM
%reinterpret_cast = memref.reinterpret_cast %arg3 to offset: [0], sizes: [8], strides: [1] : memref<?xi64> to memref<8xi64, strided<[1]>>
%alloc = memref.alloc() : memref<8xi64>
memref.copy %reinterpret_cast, %alloc : memref<8xi64, strided<[1]>> to memref<8xi64>
%0 = bufferization.to_tensor %alloc restrict writable : memref<8xi64>
hivm.hir.debug {debugtype = "print", hex = false, prefix = " x: ", tcoretype = #hivm.tcore_type<CUBE_OR_VECTOR>} %0 : tensor<8xi64>

InlineFixpipe

插入fixpipe以用于hivm.print,该hivm.print会打印mmad结果,而mmad结果是scf.for中的yield

// Before InlineFixpipe
%init = tensor.empty()
%res = scf.for iter_arg(%arg = %init) {
    %t = hivm.mmadL1 ins() outs(%arg)
    hivm.print %t
    scf.yield %t
}

// After InlineFixpipe
%init = tensor.empty()
%res = scf.for iter_arg(%arg = %init) {
    %t = hivm.mmadL1 ins() outs(%arg)
    %fixpipe = hivm.fixpipe int(%t)
    hivm.print %fixpipe
    scf.yield %t
}

InsertNZ2NDForDebug

device_print仅支持UB/GM上的数据打印,因此当打印L1的数据时,需要将数据先从L1搬至GM。该Pass的作用就是:当识别到hivm::MmadL1Op时,检查该op的输入;若输入被hivm::DebugOp用到,则需要申请一块workspace的大小,然后插入NZ2ND的op,确保搬至GM打印。

// Before InsertNZ2NDForDebug
%12 = bufferization.to_tensor %alloc restrict writable : memref<1x4xf32>
%13 = arith.index_cast %arg8 : i32 to index
%14 = arith.index_cast %5 : i32 to index
%reinterpret_cast_0 = memref.reinterpret_cast %arg4 to offset: [%14], sizes: [4, 1], strides: [%13, 1] : memref<?xf32> to memref<4x1xf32, strided<[?, 1], offset: ?>>
%alloc_1 = memref.alloc() : memref<4x1xf32>
hivm.hir.load ins(%reinterpret_cast_0 : memref<4x1xf32, strided<[?, 1], offset: ?>>) outs(%alloc_1 : memref<4x1xf32>) init_out_buffer = false may_implicit_transpose_with_last_axis = false
%15 = bufferization.to_tensor %alloc_1 restrict writable : memref<4x1xf32>
%16 = arith.muli %8, %arg8 : i32
%17 = arith.index_cast %16 : i32 to index
%18 = arith.addi %17, %14 : index
%19 = tensor.empty() : tensor<1x1xf32>
%c1 = arith.constant 1 : index
%c4 = arith.constant 4 : index
%c1_2 = arith.constant 1 : index
%20 = hivm.hir.mmadL1 {fixpipe_already_inserted = true} ins(%12, %15, %true, %c1, %c4, %c1_2 : tensor<1x4xf32>, tensor<4x1xf32>, i1, index, index, index) outs(%19 : tensor<1x1xf32>) -> tensor<1x1xf32>
hivm.hir.debug {debugtype = "print", hex = false, prefix = " a_vals: ", tcoretype = #hivm.tcore_type<CUBE_OR_VECTOR>} %12 : tensor<1x4xf32>

// After InsertNZ2NDForDebug
%12 = bufferization.to_tensor %alloc restrict writable : memref<1x4xf32>
%13 = memref_ext.alloc_workspace() : memref<1x4xf32>
%14 = bufferization.to_tensor %13 restrict writable : memref<1x4xf32>
%15 = hivm.hir.nz2nd ins(%12 : tensor<1x4xf32>) outs(%14 : tensor<1x4xf32>) -> tensor<1x4xf32>
%16 = arith.index_cast %arg8 : i32 to index
%17 = arith.index_cast %5 : i32 to index
%reinterpret_cast_0 = memref.reinterpret_cast %arg4 to offset: [%17], sizes: [4, 1], strides: [%16, 1] : memref<?xf32> to memref<4x1xf32, strided<[?, 1], offset: ?>>
%alloc_1 = memref.alloc() : memref<4x1xf32>
hivm.hir.load ins(%reinterpret_cast_0 : memref<4x1xf32, strided<[?, 1], offset: ?>>) outs(%alloc_1 : memref<4x1xf32>) init_out_buffer = false may_implicit_transpose_with_last_axis = false
%18 = bufferization.to_tensor %alloc_1 restrict writable : memref<4x1xf32>
%19 = arith.muli %8, %arg8 : i32
%20 = arith.index_cast %19 : i32 to index
%21 = arith.addi %20, %17 : index
%22 = tensor.empty() : tensor<1x1xf32>
%23 = hivm.hir.mmadL1 {fixpipe_already_inserted = true} ins(%12, %18, %true, %c1, %c4, %c1 : tensor<1x4xf32>, tensor<4x1xf32>, i1, index, index, index) outs(%22 : tensor<1x1xf32>) -> tensor<1x1xf32>
hivm.hir.debug {debugtype = "print", hex = false, prefix = " a_vals: ", tcoretype = #hivm.tcore_type<CUBE_OR_VECTOR>} %15 : tensor<1x4xf32>

SplitMixKernel

mix类用例Debug op会在该Pass内先进行InferCoreType推断出精确的coretype(VECTOR/CUBE),默认是CUBE_OR_VECTOR,然后对mix函数进行拆分后生成纯cube函数和纯vector函数,这将决定Debug op最终在cube核上运行还是vector核上运行。

InsertInitAndFinishForDebug

若存在Debug op则将hivm.hir.init_print调用添加到每个函数开头,将hivm.hir.finish_print添加到每个hivm.hir.print之后。hivm.hir.init_print用于打印之前的准备工作,hivm.hir.finish_print用于打印之后的工作,目前没有特别具体作用,为将来扩展device_print预留了接口。

// Before InsertInitAndFinishForDebug
hivm.hir.mmadL1 {fixpipe_already_inserted = true} ins(%cast, %cast_1, %true, %c1, %c4, %c1 : memref<?x?x?x?xf32, #hivm.address_space<cbuf>>, memref<?x?x?x?xf32, #hivm.address_space<cbuf>>, i1, index, index, index) outs(%cast_2 : memref<?x?x?x?xf32, #hivm.address_space<cc>>) sync_related_args(%c1_i64, %c0_i64, %c-1_i64, %c-1_i64, %c-1_i64, %c-1_i64, %c-1_i64 : i64, i64, i64, i64, i64, i64, i64)
hivm.hir.set_flag[<PIPE_M>, <PIPE_FIX>, <EVENT_ID0>]
%16 = arith.index_cast %2 : i64 to index
%17 = affine.apply affine_map<()[s0] -> (s0 * 4)>()[%16]
%view = memref.view %arg2[%17][] : memref<?xi8, #hivm.address_space<gm>> to memref<1x1xf32, #hivm.address_space<gm>>
hivm.hir.wait_flag[<PIPE_M>, <PIPE_FIX>, <EVENT_ID0>]
hivm.hir.fixpipe {enable_nz2nd} ins(%cast_2 : memref<?x?x?x?xf32, #hivm.address_space<cc>>) outs(%view : memref<1x1xf32, #hivm.address_space<gm>>)
hivm.hir.pipe_barrier[<PIPE_ALL>]
hivm.hir.sync_block_set[<CUBE>, <PIPE_FIX>, <PIPE_S>] flag = 0 ffts_base_addr = %arg0
hivm.hir.debug {debugtype = "print", hex = false, prefix = " acc_11: ", tcoretype = #hivm.tcore_type<CUBE_OR_VECTOR>} %view : memref<1x1xf32, #hivm.address_space<gm>>

// After InsertInitAndFinishForDebug
hivm.hir.init_debug
hivm.hir.mmadL1 {fixpipe_already_inserted = true} ins(%cast, %cast_1, %true, %c1, %c4, %c1 : memref<?x?x?x?xf32, #hivm.address_space<cbuf>>, memref<?x?x?x?xf32, #hivm.address_space<cbuf>>, i1, index, index, index) outs(%cast_2 : memref<?x?x?x?xf32, #hivm.address_space<cc>>) sync_related_args(%c1_i64, %c0_i64, %c-1_i64, %c-1_i64, %c-1_i64, %c-1_i64, %c-1_i64 : i64, i64, i64, i64, i64, i64, i64)
hivm.hir.set_flag[<PIPE_M>, <PIPE_FIX>, <EVENT_ID0>]
%14 = arith.index_cast %0 : i64 to index
%15 = affine.apply affine_map<()[s0] -> (s0 * 4)>()[%14]
%view = memref.view %arg2[%15][] : memref<?xi8, #hivm.address_space<gm>> to memref<1x1xf32, #hivm.address_space<gm>>
hivm.hir.wait_flag[<PIPE_M>, <PIPE_FIX>, <EVENT_ID0>]
hivm.hir.fixpipe {enable_nz2nd} ins(%cast_2 : memref<?x?x?x?xf32, #hivm.address_space<cc>>) outs(%view : memref<1x1xf32, #hivm.address_space<gm>>)
hivm.hir.pipe_barrier[<PIPE_ALL>]
hivm.hir.sync_block_set[<CUBE>, <PIPE_FIX>, <PIPE_S>] flag = 0 ffts_base_addr = %arg0
hivm.hir.debug {debugtype = "print", finishInserted = 0 : i32, hex = false, prefix = " acc_11: ", tcoretype = #hivm.tcore_type<CUBE_OR_VECTOR>} %view : memref<1x1xf32, #hivm.address_space<gm>>
hivm.hir.finish_debug

ConvertHIVMToStandard

hivm.hir.init_print/hivm.hir.print/hivm.hir.finish_print转换为库函数调用。

ConvertHIVMToLLVM

ConvertHIVMToLLVM引入真正的库函数,并设置print相关函数的链接为ExternWeak(允许在多个llvm modules中重复定义)。

Debug op库实现

op库当前实现是通过scalar打印来实现的,通过for循环外抛的方式调用毕昇编译器提供的cce::printf接口进行scalar打印。

毕昇编译器

triton-ascend产生的Host侧launcher调用bisheng编译器编好的kernel并将打印缓冲区传给kernel,待kernel返回后在Host Launcher侧读取缓冲区并进行真正的打印。此部分代码在bisheng编译器自带的头文件中实现,并由triton-ascend自动从bisheng编译器的路径中抽取。

接口说明

通过设置环境变量TRITON_DEVICE_PRINT=1来开启该功能。开启后,triton-ascend侧会设置相关宏信息__CCE_ENABLE_PRINT__,该宏信息在毕昇编译器侧会影响是否开启打印。此外,编译meta op库的时候需开启--cce-enable-print(当前默认一直开启),以确保开启打印。

// hfusion op接口
// dtype - 待打印tensor/scalar对应的数据类型
hfusion.print " prefix = xxx " {hex = xxx} %args : dtype

// hivm op接口
// tcoretype - 指明是在core核上运行还是vector核上运行(默认初始值为: CUBE_OR_VECTOR)
hivm.hir.debug {debugtype = "print", hex = xxx, prefix = " xxx: ", tcoretype = #hivm.tcore_type<CUBE_OR_VECTOR>} %args : dtype

使用约束

适用硬件

使用约束

  • Ascend 950PR/Ascend 950DT
  • Atlas A3训练系列产品/Atlas A3推理系列产品
  • Atlas A2训练系列产品/Atlas A2推理系列产品

1. 打印对象仅支持张量、标量。
2. device_print打印缓冲区固定为16KB。
3. Triton内存检测工具sanitizer与device_print互斥,不可同时启用。
4. 编码规范:单个张量单独打印,打印指令紧跟目标张量,防止张量生命周期变动引发运行异常。
5. 内核限制:不允许待打印算子仅作为device_print唯一输入。
6. 循环限制:while循环内禁止打印循环体外定义的操作数。
7. 超时限制:打印等待内核完成超时时间10分钟,长耗时用例开启打印会触发超时失败。

  • Atlas A3训练系列产品/Atlas A3推理系列产品
  • Atlas A2训练系列产品/Atlas A2推理系列产品

支持打印数据类型:boolint8uint8int16uint16int32uint32int64bfloat16halffloat32

  • Ascend 950PR/Ascend 950DT

1. 数据类型兼容:兼容Atlas A3训练系列产品/Atlas A3推理系列产品全部类型,额外支持fp8
2. 融合调度约束:插入device_print破坏VF融合边界情况下可能引发UB溢出,需减小tiling分块。
3. 缓存资源约束:打印fp8张量、L1张量边界情况下可能会引发UB溢出,需减小tiling分块规避缓存溢出。