【infra之路】MLIR中dialect的转换过程

我标注"执行了什么 Pass "以及"IR 树发生了什么变化 "。

假设我们的模型是:Y = ReLU(MatMul(A, B))。
"一个静态计算图被转换为 MLIR 后,初始是一堆高层的 Ops(如 torch dialect)。编译器运行一系列 Pass,将这些高层 Ops 逐步转换为中层 Ops(如 linalg, tensor),在这个过程中进行图优化(融合、折叠)。接着,这些中层 Ops 被进一步转换为底层 Ops(如 scf, memref, arith),在这个过程中进行循环优化和内存规划。最后,所有剩余的 Ops 都被统一转换为 llvm dialect 的 Ops,导出为 LLVM IR。"


第 0 步:初始状态(刚从前端进来,纯高层抽象)

这是我们的起点。所有东西都是高层的、声明式的。没有内存概念,没有循环概念。

mlir 复制代码
// 【状态:纯高层】
module {                    // builtin dialect
  func.func @main(%A: tensor<4x8xf32>, %B: tensor<8x4xf32>) -> tensor<4x4xf32> {
    %0 = tensor.empty() : tensor<4x4xf32>           // tensor dialect
    %1 = linalg.matmul ins(%A, %B) outs(%0) : ...   // linalg dialect (矩阵乘法)
    %2 = linalg.generic { ... } ins(%1) outs(...)   // linalg dialect (代表 ReLU)
    return %2 : tensor<4x4xf32>
  }
}

第 1 步:运行 Pass convert-linalg-to-loops(只针对 matmul)

编译器决定先把最耗时的矩阵乘法降级。它运行了一个 Pass,只把 linalg.matmul 变成了循环和内存操作 。注意,这个 Pass 不认识 tensor.empty,也不认识 ReLU 的 linalg.generic,所以它们原封不动。

变化: %1 从 linalg 变成了 scf + memref + arith。引入了内存分配 memref.alloc。

mlir 复制代码
// 【状态:混合状态 1】
module {
  func.func @main(%A: tensor<4x8xf32>, %B: tensor<8x4xf32>) -> tensor<4x4xf32> {
    %0 = tensor.empty() : tensor<4x4xf32>           // tensor dialect (没变)
    
    // --- 原来的 linalg.matmul 消失了,变成了下面这坨 ---
    %1_mem = memref.alloc() : memref<4x4xf32>       // memref dialect (分配内存)
    scf.for %i = 0 to 4 step 1 {                    // scf dialect (循环)
      scf.for %j = 0 to 4 step 1 {
        scf.for %k = 0 to 8 step 1 {
          %a = memref.load %A[%i, %k] : ...         // memref dialect
          %b = memref.load %B[%k, %j] : ...         // memref dialect
          %mul = arith.mulf %a, %b : f32            // arith dialect (标量乘)
          %add = arith.addf %sum, %mul : f32        // arith dialect (标量加)
          memref.store %add, %1_mem[%i, %j] : ...   // memref dialect
        }
      }
    }
    %1 = tensor.cast %1_mem : memref<4x4xf32> to tensor<4x4xf32> // 强行转回 tensor 以适配下游
    
    // --- ReLU 还在,因为 Pass 只处理了 matmul ---
    %2 = linalg.generic { ... } ins(%1) outs(...)   // linalg dialect (没变)
    
    return %2 : tensor<4x4xf32>
  }
}

(注:为了简化,我省略了 tensor.cast 的细节,实际中 MLIR 会用 bufferization 接口来处理这种类型不匹配,但核心意思是 %1 已经变成内存了。)


第 2 步:运行 Pass one-shot-bufferize(内存化)

现在编译器发现,IR 里既有 tensor(抽象值)又有 memref(内存),太混乱了。它运行 Bufferization Pass,把所有 tensor 都变成 memref。

变化: tensor.empty 变成了 memref.alloc。ReLU 的 linalg.generic 的输入输出也从 tensor 变成了 memref。

mlir 复制代码
// 【状态:混合状态 2】
module {
  func.func @main(%A: memref<4x8xf32>, %B: memref<8x4xf32>) -> memref<4x4xf32> {
    
    // tensor.empty 变成了 memref.alloc
    %0 = memref.alloc() : memref<4x4xf32>           // memref dialect
    
    // --- 第一步转换好的 matmul 循环 (没变) ---
    %1_mem = memref.alloc() : memref<4x4xf32>       
    scf.for %i = 0 to 4 step 1 { ... }              // scf + memref + arith
    
    // --- ReLU 的 linalg.generic 还在,但类型变了 ---
    // 它现在接收 memref,输出 memref
    %2 = linalg.generic { ... } ins(%1_mem) outs(%0) // linalg dialect (依然是高层壳子)
    
    return %0 : memref<4x4xf32>
  }
}

第 3 步:运行 Pass convert-linalg-to-loops(这次针对 ReLU)

现在轮到 ReLU 了。编译器再次运行 convert-linalg-to-loops,把剩下的 linalg.generic 也变成循环。

变化: %2 从 linalg 变成了 scf + memref + arith。

mlir 复制代码
// 【状态:纯底层混合 (全是 scf + memref + arith)】
module {
  func.func @main(%A: memref<4x8xf32>, %B: memref<8x4xf32>) -> memref<4x4xf32> {
    
    %0 = memref.alloc() : memref<4x4xf32>           // memref dialect
    
    // --- 矩阵乘法循环 (没变) ---
    %1_mem = memref.alloc() : memref<4x4xf32>       
    scf.for %i = 0 to 4 step 1 { ... }              // scf + memref + arith
    
    // --- 原来的 ReLU (linalg.generic) 消失了,变成了下面这坨 ---
    scf.for %i = 0 to 4 step 1 {                    // scf dialect (循环)
      scf.for %j = 0 to 4 step 1 {
        %val = memref.load %1_mem[%i, %j] : ...     // memref dialect (读)
        %zero = arith.constant 0.0 : f32            // arith dialect
        %relu = arith.maxf %val, %zero : f32        // arith dialect (ReLU 计算)
        memref.store %relu, %0[%i, %j] : ...        // memref dialect (写)
      }
    }
    
    return %0 : memref<4x4xf32>
  }
}

注意: 此时,所有的高级 Dialect(tensor, linalg)已经全部消失了!IR 树变成了纯粹的 scf + memref + arith。这时的代码已经非常像 C 语言了。


第 4 步:运行 Pass convert-scf-to-cf 和 convert-memref-to-llvm

现在编译器要准备交给 LLVM 了。LLVM 不认识 scf.for,也不认识 memref。它只认识"基本块(Basic Block)"和"裸指针(Pointer)"。

编译器运行两个 Pass:

  1. convert-scf-to-cf:把 scf.for 变成 cf.br(跳转指令)。
  2. convert-memref-to-llvm:把 memref 变成 llvm.ptr 和 GEP(指针计算)。

变化: scf 和 memref 消失了,变成了 cf 和 llvm。

mlir 复制代码
// 【状态:即将进入 LLVM】
module {
  // 函数签名变成了裸指针
  llvm.func @main(%A: !llvm.ptr, %B: !llvm.ptr) -> !llvm.ptr { 
    
    // memref.alloc 变成了 llvm 的堆分配调用
    %0 = llvm.call @malloc(...) : ...               // llvm dialect
    
    // --- 矩阵乘法循环变成了基本块和跳转 ---
    // (这里代码极长,简化表示)
    cf.br ^bb1(...)                                 // cf dialect (跳转)
  ^bb1:
    // 指针计算和加载
    %ptr = llvm.getelementptr ...                   // llvm dialect
    %val = llvm.load %ptr : ...                     // llvm dialect
    ...
    
    // --- ReLU 循环也变成了基本块和跳转 ---
    cf.br ^bb_reLU(...)
  ^bb_reLU:
    %relu_ptr = llvm.getelementptr ...
    %relu_val = llvm.load %relu_ptr ...
    %zero = llvm.mlir.constant(0.0) : f32           // llvm dialect
    %relu = llvm.intr.maxnum(%relu_val, %zero)      // llvm dialect (调用 LLVM 内置 max)
    llvm.store %relu, ...
    cf.br ^bb_end(...)
    
  ^bb_end:
    llvm.return %0 : !llvm.ptr
  }
}

第 5 步:运行 Pass convert-to-llvm(最终统一)

最后,编译器运行一个"大扫除"Pass,把剩下的 cf(跳转)和 arith(算术)也全部翻译成 llvm dialect。

变化: 所有非 llvm 的东西全部消失。

mlir 复制代码
// 【状态:纯 LLVM Dialect】
module {
  llvm.func @main(%A: !llvm.ptr, %B: !llvm.ptr) -> !llvm.ptr { 
    %0 = llvm.call @malloc(...) : ...
    
    // cf.br 变成了 llvm.br
    llvm.br ^bb1(...)
  ^bb1:
    %ptr = llvm.getelementptr ...
    %val = llvm.load %ptr : ...
    // arith.addf 变成了 llvm.fadd
    %add = llvm.fadd %val, %val : f32               // llvm dialect
    ...
    llvm.br ^bb_reLU(...)
  ^bb_reLU:
    // arith.maxf 变成了 llvm.intr.maxnum
    %relu = llvm.intr.maxnum(%relu_val, %zero)      // llvm dialect
    ...
    llvm.return %0 : !llvm.ptr
  }
}

第 6 步:导出为 LLVM IR

最后,MLIR 调用 mlir-translate --mlir-to-llvmir,把这棵纯 llvm dialect 的树,翻译成标准的 LLVM IR(.ll 文件)。

llvm 复制代码
; 【最终的 LLVM IR】
define ptr @main(ptr %A, ptr %B) {
entry:
  %0 = call ptr @malloc(...)
  br label %bb1

bb1:
  %ptr = getelementptr ...
  %val = load float, ptr %ptr
  %add = fadd float %val, %val
  br label %bb_relu

bb_relu:
  %relu_val = load float, ptr %relu_ptr
  %relu = call float @llvm.maxnum.f32(float %relu_val, float 0.0)
  store float %relu, ptr ...
  br label %bb_end

bb_end:
  ret ptr %0
}

总结

现在可以清晰地看到:

  1. 第 0 步: 纯 tensor + linalg(声明式,像数学公式)。
  2. 第 1-2 步: 混搭状态(tensor 和 memref 共存,linalg 和 scf 共存)。
  3. 第 3 步: 纯 scf + memref + arith(命令式,像 C 语言)。
  4. 第 4-5 步: 纯 llvm dialect(底层,像汇编的抽象)。
  5. 第 6 步: 标准 LLVM IR。

这就是 MLIR "渐进式降级"和"混合方言"的完整慢动作回放。 每一个 Pass 只做自己负责的那一小块转换,逐步把高层抽象"融化"成底层机器码。

相关推荐
梦帮科技4 小时前
端侧编译原理:TVM / MLIR 计算图模式匹配、算子融合与显存生命周期复用
网络·人工智能·深度学习·神经网络·自然语言处理·cnn·mlir
Ivanqhz19 天前
Ping-Pong 双缓冲
开发语言·人工智能·python·深度学习·mlir
Ivanqhz20 天前
BM1688 双核架构
python·深度学习·神经网络·cnn·mlir
xier_ran21 天前
【infra之路】MLIR 中 Region 和 Block 的关系
mlir
Ivanqhz21 天前
Unigram 算法
开发语言·人工智能·python·深度学习·mlir
Ivanqhz24 天前
MLIR OpBuilder
开发语言·python·mlir
清钟沁桐1 个月前
mlir 编译器学习笔记之十四 -- cpu算子自动生成
笔记·学习·mlir
早日退休!!!2 个月前
MLIR 层次化编译规范
mlir
小L~~~3 个月前
MLIR学习笔记
笔记·学习·mlir