#
Triton is a language and compiler for creating high-performance GPU kernels. AOTInductor is a PyTorch backend compiler that can generate ahead-of-time optimized deep learning inference engines and is compatible with custom Triton kernels and performance autotuning. In order to confirm if the generated Triton kernels have certain behaviors, sometimes it is necessary to inspect the SASS (Streaming Assembly) of the generated CUDA kernels.
In this blog post, I would like to demonstrate how to inspect SASS of Triton kernels from AOTInductor artifacts.
#
Based on the base address of the input tensor and the static offset for data access, the Triton compiler can deduce the appropriate vectorization instructions to generate efficient memory access patterns. In PyTorch, all the tensors allocated from PyTorch cached allocator are 256-byte aligned by default. If the stride offsets are also aligned to 16 bytes, the Triton compiler can generate 128-bit LDG.E.128 / STG.E.128 vectorized memory instructions.
In this example, we will verify if 128-bit vectorized memory instructions are generated in the SASS of CUDA kernel of interest. To verify this, because AOTInductor generates Triton kernels as cubin files, we can inspect the generated SASS to check if the 128-bit vectorized memory instructions are indeed emitted. The cubin files are saved in the pt2 zip archive and SASS can be extracted for inspection using nvdisasm or cuobjdump --dump-sass after unzipping the archive.
|
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214
|
import osimport shutilimport subprocessimport zipfileimport torchimport torch.nn as nnfrom torch.library import triton_op, wrap_tritonimport tritonimport triton.language as tl# ---------------------------------------------------------------------------# 1. Define Triton Vector Copy Kernel# ---------------------------------------------------------------------------@triton.jitdef _triton_copy_kernel( in_ptr, out_ptr, n_elements, BLOCK_SIZE: "tl.constexpr",): pid = tl.program_id(axis=0) block_start = pid * BLOCK_SIZE offsets = block_start + tl.arange(0, BLOCK_SIZE) # 16-byte alignment hint to trigger 128-bit LDG.E.128 / STG.E.128 vectorization # offsets = tl.multiple_of(offsets, 16) mask = offsets < n_elements x = tl.load(in_ptr + offsets, mask=mask) tl.store(out_ptr + offsets, x, mask=mask)# ---------------------------------------------------------------------------# 2. Transparently Traceable Triton Operator# ---------------------------------------------------------------------------@triton_op("custom_ops::triton_copy", mutates_args=())def triton_copy_op(x: torch.Tensor) -> torch.Tensor: out = torch.empty_like(x) if x.is_cuda: n_elements = x.numel() def grid(meta): return (triton.cdiv(n_elements, meta["BLOCK_SIZE"]), ) wrap_triton(_triton_copy_kernel)[grid](x, out, n_elements, BLOCK_SIZE=1024) else: out.copy_(x) return out# ---------------------------------------------------------------------------# 3. PyTorch Model# ---------------------------------------------------------------------------class VectorCopyModel(nn.Module): def forward(self, x: torch.Tensor) -> torch.Tensor: return triton_copy_op(x)# ---------------------------------------------------------------------------# 4. Main Execution Pipeline# ---------------------------------------------------------------------------def main(): device = "cuda" if torch.cuda.is_available() else "cpu" if device != "cuda": raise RuntimeError( "CUDA device required for AOTInductor Triton compilation.") print("\n" + "=" * 75) print( " PIPELINE: Export -> AOTI Compile -> Bitwise Test -> SASS Inspection") print("=" * 75) model = VectorCopyModel().to(device).eval() example_args = (torch.randn(1024 * 1024, device=device, dtype=torch.float32), ) # -------------------------------------------------------------------------- # Step 1: Export Graph # -------------------------------------------------------------------------- print( "\n[STEP 1] Exporting graph with torch.export.export(..., strict=True)" ) ep = torch.export.export(model, args=example_args, strict=True) decomposed_ep = ep.run_decompositions() print(" β
ExportedProgram graph successfully captured and decomposed.") # -------------------------------------------------------------------------- # Step 2: Compile & Package via AOTInductor # -------------------------------------------------------------------------- print("\n[STEP 2] Compiling to AOTI Package (.pt2 artifact)") output_dir = os.path.abspath("./aoti_output") os.makedirs(output_dir, exist_ok=True) pkg_path = os.path.join(output_dir, "model.pt2") compiled_pkg = torch._inductor.aoti_compile_and_package( decomposed_ep, package_path=pkg_path) print(f" β
Compiled Package Generated: {compiled_pkg}") # -------------------------------------------------------------------------- # Step 3: Run Executable & Verify Bitwise Identity # -------------------------------------------------------------------------- print("\n[STEP 3] Executing Compiled Model and Verifying Bitwise Identity") runner = torch._inductor.aoti_load_package(compiled_pkg) test_input = torch.randn(1024 * 1024, device=device, dtype=torch.float32) output = runner(test_input) assert torch.equal(test_input, output), "FAIL: Tensors are not equal!" input_bits = test_input.view(torch.int32) output_bits = output.view(torch.int32) mismatches = (input_bits != output_bits).sum().item() assert mismatches == 0, f"FAIL: Found {mismatches} bitwise mismatch(es)!" print( " β
Bitwise Identity Confirmed: Input and Output are 100% bitwise identical!" ) # -------------------------------------------------------------------------- # Step 4: Extract All Artifacts from .pt2 Package # -------------------------------------------------------------------------- print("\n[STEP 4] Unpacking .pt2 Archive and Scanning Artifacts") extract_dir = os.path.join(output_dir, "extracted") os.makedirs(extract_dir, exist_ok=True) extracted_cubins = [] extracted_so = None extracted_cpp = None with zipfile.ZipFile(compiled_pkg, "r") as zip_ref: zip_ref.extractall(extract_dir) for root, _, files in os.walk(extract_dir): for f in files: full_path = os.path.join(root, f) if f.endswith(".cubin"): extracted_cubins.append(full_path) elif f.endswith(".wrapper.so") or (f.endswith(".so") and not extracted_so): extracted_so = full_path elif f.endswith(".wrapper.cpp") or (f.endswith(".cpp") and not extracted_cpp): extracted_cpp = full_path print(f" β
Found {len(extracted_cubins)} .cubin file(s) in package.") if extracted_cpp: print(f" β
Found wrapper C++: {extracted_cpp}") if extracted_so: print(f" β
Found wrapper .so: {extracted_so}") # -------------------------------------------------------------------------- # Step 5: SASS Disassembly on Extracted .cubin via nvdisasm # -------------------------------------------------------------------------- print("\n[STEP 5] Running SASS Disassembly (nvdisasm) on Extracted .cubin") if extracted_cubins and shutil.which("nvdisasm"): for cubin_path in extracted_cubins: cubin_name = os.path.basename(cubin_path) res = subprocess.run( ["nvdisasm", "-g", cubin_path], capture_output=True, text=True, ) sass_output = res.stdout if "LDG.E.128" in sass_output or "STG.E.128" in sass_output: print( f" β
[nvdisasm CONFIRMED in {cubin_name}] Found 128-bit vector instructions:" ) for line in sass_output.splitlines(): if "LDG.E.128" in line or "STG.E.128" in line: print(f" {line.strip()}") else: print(f" --> Disassembly completed for {cubin_name}.") elif not shutil.which("nvdisasm"): print(" --> [SKIPPED] 'nvdisasm' not found in system PATH.") # -------------------------------------------------------------------------- # Step 6: SASS Disassembly on Extracted .cubin via cuobjdump # -------------------------------------------------------------------------- print( "\n[STEP 6] Running SASS Disassembly (cuobjdump) on Extracted .cubin") if extracted_cubins and shutil.which("cuobjdump"): for cubin_path in extracted_cubins: cubin_name = os.path.basename(cubin_path) res = subprocess.run( ["cuobjdump", "-sass", cubin_path], capture_output=True, text=True, ) cubin_sass = res.stdout if "LDG.E.128" in cubin_sass or "STG.E.128" in cubin_sass: print( f" β
[cuobjdump CONFIRMED in {cubin_name}] Found 128-bit vector instructions:" ) for line in cubin_sass.splitlines(): if "LDG.E.128" in line or "STG.E.128" in line: print(f" {line.strip()}") else: print(f" --> Disassembly of {cubin_name} completed.") elif not shutil.which("cuobjdump"): print(" --> [SKIPPED] 'cuobjdump' not found in system PATH.")if __name__ == "__main__": main()
|
|
12345678910111213141516171819202122232425262728293031323334353637
|
$ python triton_copy.py=========================================================================== PIPELINE: Export -> AOTI Compile -> Bitwise Test -> SASS Inspection===========================================================================[STEP 1] Exporting graph with torch.export.export(..., strict=True)/usr/lib/python3.12/copyreg.py:99: FutureWarning: `isinstance(treespec, LeafSpec)` is deprecated, use `isinstance(treespec, TreeSpec) and treespec.is_leaf()` instead. return cls.__new__(cls, *args) β
ExportedProgram graph successfully captured and decomposed.[STEP 2] Compiling to AOTI Package (.pt2 artifact)/usr/lib/python3.12/copyreg.py:99: FutureWarning: `isinstance(treespec, LeafSpec)` is deprecated, use `isinstance(treespec, TreeSpec) and treespec.is_leaf()` instead. return cls.__new__(cls, *args) β
Compiled Package Generated: /mnt/aoti_output/model.pt2[STEP 3] Executing Compiled Model and Verifying Bitwise Identity β
Bitwise Identity Confirmed: Input and Output are 100% bitwise identical![STEP 4] Unpacking .pt2 Archive and Scanning Artifacts β
Found 1 .cubin file(s) in package. β
Found wrapper C++: /mnt/aoti_output/extracted/model/data/aotinductor/model/c4si55h2cct2cbdghpljslclh43nkx45u6pefe4b6nymbkyhgoqi.wrapper.cpp β
Found wrapper .so: /mnt/aoti_output/extracted/model/data/aotinductor/model/c4si55h2cct2cbdghpljslclh43nkx45u6pefe4b6nymbkyhgoqi.wrapper.so[STEP 5] Running SASS Disassembly (nvdisasm) on Extracted .cubin β
[nvdisasm CONFIRMED in crx777x7kmdkuv7227akgtzfojh4osuapriqhz3jesxvwqpwmfnk.cubin] Found 128-bit vector instructions: /*0110*/ @!P0 LDG.E.128 R8, desc[UR4][R2.64] ; /*0140*/ @!P1 LDG.E.128 R12, desc[UR4][R2.64+0x800] ; /*0150*/ @!P0 STG.E.128 desc[UR4][R4.64], R8 ; /*0170*/ STG.E.128 desc[UR4][R4.64+0x800], R12 ;[STEP 6] Running SASS Disassembly (cuobjdump) on Extracted .cubin β
[cuobjdump CONFIRMED in crx777x7kmdkuv7227akgtzfojh4osuapriqhz3jesxvwqpwmfnk.cubin] Found 128-bit vector instructions: /*0110*/ @!P0 LDG.E.128 R8, desc[UR4][R2.64] ; /* 0x0000000402088981 */ /*0140*/ @!P1 LDG.E.128 R12, desc[UR4][R2.64+0x800] ; /* 0x00080004020c9981 */ /*0150*/ @!P0 STG.E.128 desc[UR4][R4.64], R8 ; /* 0x0000000804008986 */ /*0170*/ STG.E.128 desc[UR4][R4.64+0x800], R12 ; /* 0x0008000c04007986 */
|
#
AOTInductor Triton SASS Inspection
https://leimao.github.io/blog/AOTInductor-Triton-SASS-Inspection/