Note
Go to the end to download the full example code.
Fused Softmax
In this tutorial, you will write a fused softmax operation that is significantly faster than PyTorch’s native op for a particular class of matrices: those whose rows can fit in the GPU’s SRAM.
In doing so, you will learn about:
The benefits of kernel fusion for bandwidth-bound operations.
Reduction operators in Triton.
Motivations
Custom GPU kernels for elementwise additions are educationally valuable but won’t get you very far in practice. Let us consider instead the case of a simple (numerically stabilized) softmax operation:
import torch
import triton
import triton.language as tl
from triton.runtime import driver
DEVICE = triton.runtime.driver.active.get_active_torch_device()
def is_hip():
return triton.runtime.driver.active.get_current_target().backend == "hip"
def is_cdna():
return is_hip() and triton.runtime.driver.active.get_current_target().arch in ('gfx940', 'gfx941', 'gfx942',
'gfx90a', 'gfx908')
def naive_softmax(x):
"""Compute row-wise softmax of X using native pytorch
We subtract the maximum element in order to avoid overflows. Softmax is invariant to
this shift.
"""
# read MN elements ; write M elements
x_max = x.max(dim=1)[0]
# read MN + M elements ; write MN elements
z = x - x_max[:, None]
# read MN elements ; write MN elements
numerator = torch.exp(z)
# read MN elements ; write M elements
denominator = numerator.sum(dim=1)
# read MN + M elements ; write MN elements
ret = numerator / denominator[:, None]
# in total: read 5MN + 2M elements ; wrote 3MN + 2M elements
return ret
When implemented naively in PyTorch, computing y = naive_softmax(x) for \(x \in R^{M \times N}\)
requires reading \(5MN + 2M\) elements from DRAM and writing back \(3MN + 2M\) elements.
This is obviously wasteful; we’d prefer to have a custom “fused” kernel that only reads
X once and does all the necessary computations on-chip.
Doing so would require reading and writing back only \(MN\) bytes, so we could
expect a theoretical speed-up of ~4x (i.e., \((8MN + 4M) / 2MN\)).
The torch.jit.script flags aims to perform this kind of “kernel fusion” automatically
but, as we will see later, it is still far from ideal.
Compute Kernel
Our softmax kernel works as follows: each program loads a set of rows of the input matrix X strided by number of programs, normalizes it and writes back the result to the output Y.
Note that one important limitation of Triton is that each block must have a power-of-two number of elements, so we need to internally “pad” each row and guard the memory operations properly if we want to handle any possible input shapes:
@triton.jit
def softmax_kernel(output_ptr, input_ptr, input_row_stride, output_row_stride, n_rows, n_cols, BLOCK_SIZE: tl.constexpr,
num_stages: tl.constexpr):
# starting row of the program
row_start = tl.program_id(0)
row_step = tl.num_programs(0)
for row_idx in tl.range(row_start, n_rows, row_step, num_stages=num_stages):
# The stride represents how much we need to increase the pointer to advance 1 row
row_start_ptr = input_ptr + row_idx * input_row_stride
# The block size is the next power of two greater than n_cols, so we can fit each
# row in a single block
col_offsets = tl.arange(0, BLOCK_SIZE)
input_ptrs = row_start_ptr + col_offsets
# Load the row into SRAM, using a mask since BLOCK_SIZE may be > than n_cols
mask = col_offsets < n_cols
row = tl.load(input_ptrs, mask=mask, other=-float('inf'))
# Subtract maximum for numerical stability
row_minus_max = row - tl.max(row, axis=0)
# Note that exponentiation in Triton is fast but approximate (i.e., think __expf in CUDA)
numerator = tl.exp(row_minus_max)
denominator = tl.sum(numerator, axis=0)
softmax_output = numerator / denominator
# Write back output to DRAM
output_row_start_ptr = output_ptr + row_idx * output_row_stride
output_ptrs = output_row_start_ptr + col_offsets
tl.store(output_ptrs, softmax_output, mask=mask)
We can create a helper function that enqueues the kernel and its (meta-)arguments for any given input tensor.
properties = driver.active.utils.get_device_properties(DEVICE.index)
NUM_SM = properties["multiprocessor_count"]
NUM_REGS = properties["max_num_regs"]
SIZE_SMEM = properties["max_shared_mem"]
WARP_SIZE = properties["warpSize"]
target = triton.runtime.driver.active.get_current_target()
kernels = {}
def softmax(x):
n_rows, n_cols = x.shape
# The block size of each loop iteration is the smallest power of two greater than the number of columns in `x`
BLOCK_SIZE = triton.next_power_of_2(n_cols)
# Another trick we can use is to ask the compiler to use more threads per row by
# increasing the number of warps (`num_warps`) over which each row is distributed.
# You will see in the next tutorial how to auto-tune this value in a more natural
# way so you don't have to come up with manual heuristics yourself.
num_warps = 8
# Number of software pipelining stages.
num_stages = 4 if SIZE_SMEM > 200000 else 2
# Allocate output
y = torch.empty_like(x)
# pre-compile kernel to get register usage and compute thread occupancy.
kernel = softmax_kernel.warmup(y, x, x.stride(0), y.stride(0), n_rows, n_cols, BLOCK_SIZE=BLOCK_SIZE,
num_stages=num_stages, num_warps=num_warps, grid=(1, ))
kernel._init_handles()
n_regs = kernel.n_regs
size_smem = kernel.metadata.shared
if is_hip():
# NUM_REGS represents the number of regular purpose registers. On CDNA architectures this is half of all registers available.
# However, this is not always the case. In most cases all registers can be used as regular purpose registers.
# ISA SECTION (3.6.4 for CDNA3)
# VGPRs are allocated out of two pools: regular VGPRs and accumulation VGPRs. Accumulation VGPRs are used
# with matrix VALU instructions, and can also be loaded directly from memory. A wave may have up to 512 total
# VGPRs, 256 of each type. When a wave has fewer than 512 total VGPRs, the number of each type is flexible - it is
# not required to be equal numbers of both types.
NUM_GPRS = NUM_REGS
if is_cdna():
NUM_GPRS = NUM_REGS * 2
# MAX_NUM_THREADS represents maximum number of resident threads per multi-processor.
# When we divide this number with WARP_SIZE we get maximum number of waves that can
# execute on a CU (multi-processor) in parallel.
MAX_NUM_THREADS = properties["max_threads_per_sm"]
max_num_waves = MAX_NUM_THREADS // WARP_SIZE
occupancy = min(NUM_GPRS // WARP_SIZE // n_regs, max_num_waves) // num_warps
else:
occupancy = NUM_REGS // (n_regs * WARP_SIZE * num_warps)
occupancy = min(occupancy, SIZE_SMEM // size_smem)
num_programs = NUM_SM * occupancy
num_programs = min(num_programs, n_rows)
# Create a number of persistent programs.
kernel[(num_programs, 1, 1)](y, x, x.stride(0), y.stride(0), n_rows, n_cols, BLOCK_SIZE, num_stages)
return y
Unit Test
We make sure that we test our kernel on a matrix with an irregular number of rows and columns. This will allow us to verify that our padding mechanism works.
torch.manual_seed(0)
x = torch.randn(1823, 781, device=DEVICE)
y_triton = softmax(x)
y_torch = torch.softmax(x, axis=1)
assert torch.allclose(y_triton, y_torch), (y_triton, y_torch)
As expected, the results are identical.
Benchmark
Here we will benchmark our operation as a function of the number of columns in the input matrix – assuming 4096 rows.
We will then compare its performance against (1) torch.softmax and (2) the naive_softmax defined above.
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=['N'], # argument names to use as an x-axis for the plot
x_vals=[128 * i for i in range(2, 100)], # different possible values for `x_name`
line_arg='provider', # argument name whose value corresponds to a different line in the plot
line_vals=['triton', 'torch', 'naive_softmax'], # possible values for `line_arg``
line_names=["Triton", "Torch", "Naive Softmax"], # label name for the lines
styles=[('blue', '-'), ('green', '-'), ('red', '-')], # line styles
ylabel="GB/s", # label name for the y-axis
plot_name="softmax-performance", # name for the plot. Used also as a file name for saving the plot.
args={'M': 4096}, # values for function arguments not in `x_names` and `y_name`
))
def benchmark(M, N, provider):
x = torch.randn(M, N, device=DEVICE, dtype=torch.float32)
stream = getattr(torch, DEVICE.type).Stream()
getattr(torch, DEVICE.type).set_stream(stream)
if provider == 'torch':
ms = triton.testing.do_bench(lambda: torch.softmax(x, axis=-1))
if provider == 'triton':
ms = triton.testing.do_bench(lambda: softmax(x))
if provider == 'naive_softmax':
ms = triton.testing.do_bench(lambda: naive_softmax(x))
gbps = lambda ms: 2 * x.numel() * x.element_size() * 1e-9 / (ms * 1e-3)
return gbps(ms)
benchmark.run(show_plots=True, print_data=True)

softmax-performance:
N Triton (GB/s) Torch (GB/s) Naive Softmax (GB/s)
0 256.0 507.030164 706.952743 206.390610
1 384.0 697.592621 817.049870 263.642738
2 512.0 828.475036 931.192685 302.948292
3 640.0 848.212190 926.920613 331.096537
4 768.0 911.227012 993.022213 350.802448
5 896.0 972.079050 1029.956144 354.562274
6 1024.0 1027.744061 1063.496708 354.487012
7 1152.0 1113.167062 1075.117175 348.731411
8 1280.0 1135.306623 1113.690743 348.413387
9 1408.0 1167.448301 1135.483899 342.744987
10 1536.0 1194.987061 1160.520820 333.679246
11 1664.0 1223.320662 1187.344009 329.936789
12 1792.0 1242.971212 1195.606596 326.388797
13 1920.0 1269.948546 1227.398765 325.792646
14 2048.0 1281.559479 1253.895138 324.976200
15 2176.0 1239.757342 964.131876 325.655362
16 2304.0 1263.037214 1000.085152 326.040586
17 2432.0 1283.185124 1038.857776 327.697690
18 2560.0 1291.719234 1068.060624 328.488138
19 2688.0 1300.820607 1096.771609 329.806882
20 2816.0 1314.635537 1123.316938 329.136083
21 2944.0 1318.855030 1144.696894 331.818460
22 3072.0 1329.571784 1174.514523 333.777652
23 3200.0 1338.219109 1172.265608 334.773944
24 3328.0 1351.028564 1196.953179 336.442669
25 3456.0 1359.118824 1220.798682 336.972753
26 3584.0 1361.894216 1244.456692 338.831140
27 3712.0 1378.660793 1267.393270 340.636443
28 3840.0 1377.431001 1283.514760 340.356898
29 3968.0 1378.981131 1296.937229 341.276321
30 4096.0 1388.271708 1319.516815 339.659762
31 4224.0 1341.708415 1278.101849 342.791246
32 4352.0 1355.376403 1300.081427 345.180280
33 4480.0 1361.322081 1318.591509 345.733399
34 4608.0 1369.242435 1334.945450 347.209349
35 4736.0 1367.748724 1341.887823 348.425254
36 4864.0 1381.780631 1361.537623 349.020783
37 4992.0 1381.197485 1371.869298 350.398995
38 5120.0 1392.157160 1383.366034 350.961178
39 5248.0 1387.220434 1359.496435 351.790523
40 5376.0 1390.805965 1373.041809 351.894830
41 5504.0 1393.622996 1376.061611 353.727128
42 5632.0 1403.150102 1387.020841 353.352493
43 5760.0 1407.739325 1405.256105 355.056686
44 5888.0 1394.048654 1420.935167 355.064568
45 6016.0 1407.264409 1417.763300 356.915069
46 6144.0 1421.858980 1431.735767 357.752505
47 6272.0 1416.475250 1394.618663 357.938114
48 6400.0 1424.492022 1402.830634 358.580619
49 6528.0 1427.569883 1416.969063 359.015301
50 6656.0 1421.651278 1431.530516 359.200009
51 6784.0 1427.555286 1439.581574 360.302832
52 6912.0 1434.139063 1447.851217 360.652116
53 7040.0 1424.020236 1453.087792 361.244442
54 7168.0 1435.754924 1461.116909 361.897677
55 7296.0 1432.898761 1089.112686 362.731069
56 7424.0 1442.341970 1099.532543 363.079471
57 7552.0 1437.020861 1110.899497 363.598815
58 7680.0 1442.428649 1124.408873 363.310930
59 7808.0 1432.562222 1134.385974 364.807906
60 7936.0 1443.717842 1139.730998 364.796750
61 8064.0 1443.675482 1150.490533 364.885971
62 8192.0 1436.890436 1149.281816 363.818189
63 8320.0 1391.761170 1116.423536 361.457857
64 8448.0 1395.485655 1123.926315 362.037866
65 8576.0 1396.505934 1126.898038 363.339047
66 8704.0 1395.997805 1134.044835 364.569904
67 8832.0 1398.523250 1133.083204 364.796027
68 8960.0 1393.755614 1140.106698 365.874289
69 9088.0 1406.235657 1135.339480 366.666634
70 9216.0 1416.861625 1142.619396 367.666765
71 9344.0 1400.557105 1421.095308 367.873119
72 9472.0 1415.140491 1432.782242 368.619779
73 9600.0 1407.005722 1428.264090 368.207078
74 9728.0 1410.724209 1439.219470 369.876066
75 9856.0 1414.339793 1439.756358 369.186688
76 9984.0 1402.313241 1451.371088 370.919642
77 10112.0 1418.444092 1454.505454 371.064839
78 10240.0 1424.269444 1465.938350 370.820033
79 10368.0 1426.916826 1464.343946 369.704500
80 10496.0 1420.552997 1464.968604 370.546270
81 10624.0 1418.918070 1466.612780 370.184847
82 10752.0 1410.245298 1472.558947 371.100751
83 10880.0 1404.548646 1480.421712 371.917247
84 11008.0 1431.093288 1474.748850 371.935129
85 11136.0 1439.100891 1484.090406 373.005865
86 11264.0 1422.871657 1487.823444 373.298920
87 11392.0 1435.064572 1490.762457 373.772309
88 11520.0 1421.869408 1491.964172 373.794889
89 11648.0 1427.961389 1500.924833 374.316934
90 11776.0 1446.362228 1502.083434 374.859530
91 11904.0 1445.123004 1510.824788 375.179709
92 12032.0 1435.836185 1509.698186 375.778932
93 12160.0 1429.672022 1515.698651 375.903782
94 12288.0 1444.022873 1417.793595 376.012273
95 12416.0 1445.803197 1396.261964 374.196037
96 12544.0 1456.758092 1392.710416 375.257260
97 12672.0 1452.483708 1390.028559 375.198421
- In the above plot, we can see that:
Triton is 4x faster than the Torch JIT. This confirms our suspicions that the Torch JIT does not do any fusion here.
Triton is noticeably faster than
torch.softmax– in addition to being easier to read, understand and maintain. Note however that the PyTorch softmax operation is more general and will work on tensors of any shape.
Total running time of the script: (0 minutes 34.695 seconds)