Enceladus v0.1.0 · alpha
Porting from TritonAll pages
OverviewQuickstartProgramming modelMemory and synchronizationLanguage referenceDebuggingFramework interopPerformanceBenchmarksPorting from Triton

Guides

Porting from Triton

Most Triton kernels port by changing the imports and the device. This page lists the steps, what differs, and the Triton features that Enceladus refuses.

Port a kernel

To port a kernel, complete the following steps:

  1. Replace import triton.language as tl with import enceladus.language as tl, and @triton.jit with @enceladus.jit.
  2. Move tensors from cuda to mps.
  3. Mask the M and N edges of the loads, because Enceladus doesn't check bounds on pointer accesses.
  4. Replace triton.cdiv with enceladus.cdiv. Hints such as num_stages can stay; Enceladus ignores them.
  5. Optional: Switch matmul loads to tensor descriptors, passing only the row strides.

What differs

The following table lists what differs:

AreaTritonEnceladus
num_warpsWarpsSIMD groups of 32 threads, a power of two from 1 to 32
ProgramCTAMetal threadgroup with 32 KB of threadgroup memory
Typesfp64, fp8, tf32No float64, FP8, or TF32. A float32 dot runs in full float32.
tl.dotMany dtypes, including integersSame-dtype float16, bfloat16, or float32 operands. K must be a multiple of 8.
DescriptorsNeeds a host allocatorNo allocator. The innermost stride must be 1.
AtomicsFull sem and scopeRelaxed ordering only. No 16-bit float atomic_add; use a float32 buffer.
New launch optionsNonedot_warps=(WM, WN) and dot_backend
Inline assemblytl.inline_asm_elementwiseenceladus.metal_kernel(source, name) for hand-written MSL
Benchmarkingtriton.testing.do_benchenceladus.testing.do_bench

Unsupported features

The following Triton features raise CompilationError:

  • Block pointers: tl.make_block_ptr, tl.advance, boundary_check, and padding_option.
  • while, break, and continue.
  • tl.multiple_of, tl.max_contiguous, tl.join, tl.split, tl.dot_scaled, tl.sort, tl.flip, tl.gather, tl.histogram, tl.rand, and libdevice.
  • ** on runtime values. Use tl.exp2 and tl.log2, or multiply.
  • Warp specialization, TMA, and clusters.