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)
02 fused softmax
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)

Gallery generated by Sphinx-Gallery