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 509.758600 702.398000 203.859411
1 384.0 708.599167 830.856386 262.126417
2 512.0 826.715611 922.522522 302.104264
3 640.0 846.059014 924.975116 332.283618
4 768.0 909.157010 991.176263 352.292260
5 896.0 973.082876 1022.570369 355.228326
6 1024.0 1018.770526 1073.689829 355.185636
7 1152.0 1034.102906 1063.479626 348.409217
8 1280.0 1066.435687 1111.711020 347.063592
9 1408.0 1112.309346 1136.228622 342.307637
10 1536.0 1146.434784 1167.073299 332.658237
11 1664.0 1183.632702 1186.339147 328.581523
12 1792.0 1199.883541 1190.866018 324.686144
13 1920.0 1237.487007 1226.040782 325.914198
14 2048.0 1249.171432 1252.280266 325.811756
15 2176.0 1178.818109 961.559002 326.370376
16 2304.0 1197.252045 1006.070304 326.592875
17 2432.0 1226.254775 1038.906888 327.189706
18 2560.0 1240.937628 1067.361606 328.685046
19 2688.0 1263.043064 1095.150193 329.404753
20 2816.0 1271.949315 1126.605114 328.503042
21 2944.0 1299.311106 1144.561854 331.650940
22 3072.0 1304.252275 1173.913156 333.384207
23 3200.0 1323.630816 1174.598243 334.925325
24 3328.0 1322.089478 1199.359701 336.161614
25 3456.0 1339.954008 1222.235543 337.224005
26 3584.0 1344.606046 1243.536689 338.570421
27 3712.0 1345.415632 1259.982015 340.928328
28 3840.0 1361.881446 1284.598696 340.865227
29 3968.0 1365.845167 1296.588372 340.696450
30 4096.0 1364.917377 1315.397786 338.501344
31 4224.0 1338.017735 1277.764199 343.604636
32 4352.0 1349.507956 1297.361753 344.864400
33 4480.0 1352.818523 1318.401925 345.616360
34 4608.0 1365.012000 1330.706631 347.418851
35 4736.0 1361.476444 1342.389073 348.660938
36 4864.0 1379.159474 1357.122112 349.239302
37 4992.0 1374.878934 1369.646181 350.340417
38 5120.0 1385.994532 1386.771818 351.409125
39 5248.0 1386.106603 1356.182156 351.931176
40 5376.0 1382.018591 1371.038953 351.870837
41 5504.0 1388.752487 1381.553564 353.539041
42 5632.0 1400.923959 1396.003124 353.478076
43 5760.0 1404.696256 1404.839022 354.716766
44 5888.0 1392.268220 1408.229167 354.468274
45 6016.0 1405.738811 1416.874722 357.076481
46 6144.0 1417.781241 1422.173093 357.524810
47 6272.0 1416.249021 1406.489823 357.946085
48 6400.0 1414.276118 1414.973860 358.826106
49 6528.0 1423.482094 1417.885610 359.301012
50 6656.0 1416.090205 1422.442155 359.130835
51 6784.0 1425.583233 1441.535083 359.935210
52 6912.0 1428.549979 1445.135712 360.498508
53 7040.0 1427.512997 1449.240461 360.988910
54 7168.0 1428.699370 1467.798323 362.001920
55 7296.0 1432.163014 1088.804760 362.740233
56 7424.0 1436.169422 1099.061647 362.835453
57 7552.0 1436.368827 1107.738998 364.070605
58 7680.0 1437.429770 1124.536395 363.778317
59 7808.0 1435.288336 1134.219286 364.530412
60 7936.0 1438.346405 1143.749595 364.932834
61 8064.0 1434.495122 1152.374508 365.379967
62 8192.0 1431.343194 1153.790961 363.610928
63 8320.0 1382.352527 1115.259572 361.569420
64 8448.0 1386.654509 1126.074445 362.104601
65 8576.0 1391.042213 1128.595079 363.450435
66 8704.0 1387.358584 1135.571017 364.315877
67 8832.0 1392.555522 1130.986086 365.080541
68 8960.0 1386.147033 1137.739987 365.705402
69 9088.0 1398.860281 1134.636259 366.777797
70 9216.0 1408.862218 1141.948488 367.725152
71 9344.0 1394.083237 1424.987183 367.631293
72 9472.0 1403.090540 1432.966173 368.546867
73 9600.0 1395.457499 1431.073903 368.832240
74 9728.0 1397.204322 1438.433375 369.768469
75 9856.0 1397.878166 1436.298859 369.961290
76 9984.0 1395.067954 1450.569821 370.282239
77 10112.0 1411.565342 1452.045197 371.387427
78 10240.0 1408.442426 1467.668840 371.896542
79 10368.0 1416.828533 1463.674770 370.108050
80 10496.0 1408.392399 1464.755578 370.213309
81 10624.0 1407.696813 1468.324213 371.063978
82 10752.0 1399.244443 1470.916587 371.237791
83 10880.0 1400.351027 1478.411918 371.426311
84 11008.0 1421.770045 1478.773846 371.622105
85 11136.0 1423.584811 1483.488086 372.460936
86 11264.0 1412.404086 1484.241307 372.739845
87 11392.0 1419.415510 1487.217528 373.657815
88 11520.0 1407.974040 1497.181278 373.671530
89 11648.0 1420.018476 1500.094577 375.048061
90 11776.0 1434.915993 1500.147370 375.023855
91 11904.0 1435.090267 1510.628412 375.382172
92 12032.0 1423.716940 1510.651094 376.154905
93 12160.0 1419.522380 1515.143782 375.726725
94 12288.0 1430.807164 1419.872706 375.994519
95 12416.0 1432.489015 1397.286285 374.652460
96 12544.0 1442.851506 1390.919904 374.928663
97 12672.0 1436.618073 1391.127023 375.084204
- 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 35.017 seconds)