Vx fully embraces {ML, IR, MLIR}

TL;DR — {ML, IR, MLIR} is shell brace expansion for the three things this post is about. ML: tensors are built-in values, a @ b becomes a linalg.matmul, and the standard library has softmax and reductions; grad differentiates scalar functions through Enzyme. IR: Vx has its own flat intermediate representation, and a library can ship as that IR instead of as source. MLIR: Vx defines its own MLIR dialect, written in TableGen, that keeps where data lives and where code runs; its own passes lower that into the upstream gpu, linalg and llvm dialects, and MLIR's NVVM path turns a GPU region into PTX. The last section lists what “fully” does not cover yet.

Three words, three meanings. ML is machine learning: the tensor programs Vx is built to compile. An IR, or intermediate representation, is the form a compiler holds a program in between reading the source and writing machine code. MLIR is the LLVM project's framework for building IRs out of dialects: sets of operations at one level of abstraction, such as linear algebra, GPU kernels or plain LLVM instructions. A compiler lowers a program by rewriting it from higher dialects into lower ones until only machine code is left.

This post follows small Vx programs down that path and shows the output of each stage. Every program here was compiled and run with vxc built from commit 411f13c7.

ML: tensors the compiler understands

A tensor's element type and shape are part of its type, so the compiler knows both before the program runs. Matrix multiplication is the @ operator:

fn main() -> i32 {
  let mut a = Tensor<f32, [2, 3]>::uninit();
  let mut b = Tensor<f32, [3, 2]>::uninit();
  for i in 0..2 {
    for k in 0..3 {
      a[i][k] = (i * 3 + k + 1) as f32;
      b[k][i] = 1.0;
    }
  }
  let c = a @ b;
  print(c[1][0]);
  return 0;
}

This prints 15: row 1 of a is 4, 5, 6, and every column of b is ones. The @ does not become three nested loops in the compiler's output. It becomes one operation from MLIR's linalg dialect, which still says “this is a matrix product of a 2×3 and a 3×2” (vxc matmul.vx --action emit-mlir):

linalg.fill ins(%cst_10 : f32) outs(%alloc_9 : memref<2x2xf32>)
linalg.matmul ins(%alloc, %alloc_3 : memref<2x3xf32>, memref<3x2xf32>) outs(%alloc_9 : memref<2x2xf32>)

Because the product is still a linalg.matmul at this point, later stages can choose how to run it. When such a product is placed on an NVIDIA GPU, it is left at this level on purpose so the runtime can hand it to cuBLAS instead of running a kernel Vx generated.

The same works for f16 and bf16. A test in the suite adds two f16 values, stores into an f16 tensor, and multiplies two f16 matrices, printing 4 1.5 6.

The standard library's std::tensor has the reductions a model needs: sum, dot, mean, max, argmax, variance, softmax_inplace and more. Here is a softmax over scores large enough that a careless exp would overflow:

import std::tensor;

fn main() -> i32 {
  // Scores for three classes, large enough that exp(1002) overflows an f64.
  let mut logits = Tensor<f64>([3]);
  logits[0] = 1000.0;
  logits[1] = 1001.0;
  logits[2] = 1002.0;

  let _done = logits.softmax_inplace();
  for i in 0..3 {
    print((logits[i] * 1000.0).round() / 1000.0);
    print!(" ");
  }
  match logits.argmax() {
    Option<i32>::Some(best) => {
      print!("best class: ");
      print(best);
    }
    Option<i32>::None => {
      print!("empty");
    }
  }
  return 0;
}

It prints 0.09 0.245 0.665 best class: 2, the same as a numerically stable softmax in Python. Parts of std::tensor are written in MLIR directly, inside Vx, and softmax_inplace is one of them: its body includes a block written with the upstream scf and vector dialects, and those operations appear in this program's MLIR as they were written. More on that below.

Training needs derivatives, and grad gives them:

fn poly(x : f32) -> f32 {
  return x * x * x;
}

fn main() -> i32 {
  let x : f32 = 2.0;
  // The derivative of x^3 is 3x^2, which is 12 at x = 2.
  let dx : f32 = grad(poly, x);
  print(dx);
  return 0;
}

The compiler turns grad(poly, x) into a call that the Enzyme automatic-differentiation plugin fills in later, on the LLVM IR:

func.func private @__enzyme_autodiff_grad_poly((f32) -> f32, f32) -> f32
%0 = call @__enzyme_autodiff_grad_poly(%f, %cst) : ((f32) -> f32, f32) -> f32

jvp and vjp, forward and reverse mode, work the same way. The continuous-integration machines build Enzyme and run these programs; the machine this post was written on does not have Enzyme, so the call above is as far as it was checked here.

IR: Vx's own, and a library can ship as it

Between the source and MLIR, Vx keeps the program in its own IR: the flat HIR, a list of simple instructions in which each value is named by its position. It is produced after type checking and borrow checking, so everything in it is already known to be well typed, and it is what the default code generator reads.

That IR is also a file format. --action emit-interface writes a module's function signatures and its flat-HIR bodies to a .vxlib file, and a later compile can import that file without the source:

fn relu(x : f32) -> f32 {
  if x > 0.0 {
    return x;
  }
  return 0.0;
}
$ vxc activations.vx --action emit-interface -o activations.vxlib
Wrote module interface (782 bytes) to activations.vxlib
  fn_sigs 1/1  methods 0/0  bodies 1/1  structs 0/0  (nothing skipped)
$ rm activations.vx
import activations;

fn main() -> i32 {
  print(relu(-2.5) + relu(4.0));
  return 0;
}

With activations.vx deleted, this program still type-checks against the signatures in the .vxlib, links the body of relu from it, and prints 4.

MLIR: a dialect of its own

The upstream dialects describe computation well, but none of them says “this buffer lives in the GPU's memory” or “run this region on that device”. Vx adds a dialect, vx, that does. It is defined in TableGen, MLIR's language for declaring operations, in include/VxDialect.td. The two central operations, with their comments and printing formats trimmed:

def Vx_SpawnOp : Vx_Op<"spawn"> {
    let summary = "Spawns execution on a specific topology.";
    let arguments = (ins I32Attr:$topology);
    let results = (outs Variadic<AnyType>:$results);
    let regions = (region AnyRegion:$body);
}

def Vx_TransferOp : Vx_Op<"transfer", [MemoryEffects<[MemWrite]>]> {
    let summary = "Transfers a tensor to a target topology.";
    let arguments = (ins Arg<AnyType, "the value to copy", [MemRead]>:$src,
                         I32Attr:$target_topology);
    let results = (outs Res<AnyType, "the copy on the target topology",
                            [MemAlloc]>:$dst);
}

The dialect also has vx.free, vx.launch, vx.kernel, vx.barrier and a few more. Here is a program that uses the first two. It copies two vectors to the GPU, computes y = 2x + y there, and copies the result back:

fn main() -> i32 {
  let mut x_host = Tensor<f32, [8]>::uninit();
  let mut y_host = Tensor<f32, [8]>::uninit();
  for i in 0..8 {
    x_host[i] = i as f32;
    y_host[i] = 1.0;
  }

  let x = transfer(x_host, Memory::GPU_HBM);
  let mut y = transfer(y_host, Memory::GPU_HBM);

  spawn on(Topology::GPU) {
    for i in 0..8 {
      y[i] = 2.0 * x[i] + y[i];
    }
  }

  let result = transfer(y, Memory::CPU_DRAM);
  print(result);
  return 0;
}

It prints [1, 3, 5, 7, 9, 11, 13, 15]. Leave out the copy back, and the type checker stops the build, because the CPU cannot read GPU memory:

Error[E6003] at 18:9: 'y' lives in GPU_HBM but CPU sees only [CPU_DRAM, NPU_HBM]; insert an explicit transfer to CPU_DRAM (cost 50 on the declared path)
Error[E6003] at 18:3: `print` reads its argument on CPU, which sees only [CPU_DRAM, NPU_HBM], but the value lives in GPU_HBM; bring it home first with `transfer(.., Memory::CPU_DRAM)`

Stage 1: the vx dialect

In the MLIR that vxc --action emit-mlir prints, the placement is still there, as operations rather than as a guess. Topology 500 is the GPU and 0 is the CPU; the loop body is left out here:

%11 = "vx.transfer"(%alloc) <{target_topology = 500 : i32}> : (memref<8xf32>) -> memref<8xf32>
%12 = "vx.transfer"(%alloc_1) <{target_topology = 500 : i32}> : (memref<8xf32>) -> memref<8xf32>
vx.spawn topology(500) {
  // the loop, with its start, bound and step marked as parallel
  vx.yield
} {vx_parallel_trip = 8 : i64}
%13 = "vx.transfer"(%12) <{target_topology = 0 : i32}> : (memref<8xf32>) -> memref<8xf32>

Stage 2: Vx's passes, into upstream dialects

Vx's own pass, convert-vx-to-standard, takes the region out into a kernel, launches it, and frees each copied buffer after its last use. The kernel body is now written in MLIR's standard gpu dialect, with each thread computing its index from gpu.block_id and gpu.thread_id:

%11 = "vx.transfer"(%alloc) <{target_topology = 500 : i32}> : (memref<8xf32>) -> memref<8xf32>
%12 = "vx.transfer"(%alloc_1) <{target_topology = 500 : i32}> : (memref<8xf32>) -> memref<8xf32>
vx.launch @vx_npu_kernel_0(%11, %12) {topology = 500 : i32, vx_parallel_trip = 8 : i64} : (memref<8xf32>, memref<8xf32>) -> ()
vx.free %11 {topology = 500 : i32} : memref<8xf32>
%13 = "vx.transfer"(%12) <{target_topology = 0 : i32}> : (memref<8xf32>) -> memref<8xf32>
vx.free %12 {topology = 500 : i32} : memref<8xf32>
gpu.module @vx_kernels {
  gpu.func @vx_npu_kernel_0(%arg0: memref<8xf32>, %arg1: memref<8xf32>) kernel {
    %c0_i32 = arith.constant 0 : i32
    %c8_i32 = arith.constant 8 : i32
    %cst = arith.constant 2.000000e+00 : f32
    %c1_i32 = arith.constant 1 : i32
    %alloca = memref.alloca() : memref<i32>
    %block_dim_x = gpu.block_dim  x
    %block_id_x = gpu.block_id  x
    %thread_id_x = gpu.thread_id  x
    %0 = arith.muli %block_id_x, %block_dim_x : index
    %1 = arith.addi %0, %thread_id_x : index
    %2 = arith.index_cast %1 : index to i32
    ...

From here on, the work is done by upstream MLIR. The GPU module goes through convert-gpu-to-nvvm and gpu-module-to-binary; the host code goes through the usual conversions to the llvm dialect, and then to LLVM IR.

Stage 3: PTX

PTX is NVIDIA's assembly language for GPUs. In the LLVM IR that vxc --action emit-llvm prints, the dispatch call carries the kernel as PTX text. Decoded, its loop is a grid-stride loop, and 2.0 * x has become x + x:

.version 7.6
.target sm_80
.address_size 64
...
	mov.u32 	%r1, %ntid.x;
	mov.u32 	%r3, %ctaid.x;
	mov.u32 	%r4, %tid.x;
	mad.lo.s32 	%r10, %r3, %r1, %r4;
	setp.gt.s32 	%p1, %r10, 7;
	@%p1 bra 	$L__BB0_3;
	mov.u32 	%r5, %nctaid.x;
	mul.lo.s32 	%r2, %r5, %r1;
$L__BB0_2:
	mul.wide.s32 	%rd5, %r10, 4;
	add.s64 	%rd6, %rd1, %rd5;
	ld.global.b32 	%r6, [%rd6];
	add.rn.f32 	%r7, %r6, %r6;
	add.s64 	%rd7, %rd2, %rd5;
	ld.global.b32 	%r8, [%rd7];
	add.rn.f32 	%r9, %r7, %r8;
	st.global.b32 	[%rd7], %r9;
	add.s32 	%r10, %r10, %r2;
	setp.lt.s32 	%p2, %r10, 8;
	@%p2 bra 	$L__BB0_2;

The test suite checks that this PTX is produced and holds the loop, not just a stub. It has no GPU to run it on. Neither had the machine this post was written on, so the runtime ran the same region on the CPU, which is how the program printed its result above.

MLIR inside Vx source

The traffic goes both ways. A Vx function can contain MLIR, through mlir!. This is fill from std::tensor:

fn fill(self : &mut Tensor<T, [?, ?]>, val : T) -> void {
  mlir!(
  inputs: (%arg_t = self: memref<?x?xf32>, %arg_v = val: f32),
  clobbers: [self],
  dialects: ["linalg"]
  ) {
    linalg.fill ins(%arg_v : f32) outs(%arg_t : memref<?x?xf32>)
    macro.yield
  };
}

The block names its inputs, the dialects it uses, and the arguments it writes. Because of clobbers, the borrow checker and the compile-time interpreter both know that fill changes self, even though they cannot read the MLIR inside.

What “fully” does not cover yet

So “fully” describes the direction more than the finish line. But the direction is concrete. Placement lives in the types, survives into Vx's own IR, is spelled out as operations in Vx's own MLIR dialect, and is handed to upstream MLIR only at the point where it turns into kernels and copies. Nothing along the way has to guess where a value lives.

For why keeping structure until code generation matters, see Premature de-optimization is the root of all evil. For GPU tiles in shared memory, see Tiles, without a tile type. The dialect is in include/VxDialect.td and its lowering in src/dialect/VxLowering.cpp.