PyTorch has added a FlatStrided execution plan for eligible full reductions on non-contiguous MPS tensors, letting supported workloads reduce strided views in place instead of materializing a contiguous copy first. For Apple GPU users running operations such as sum or max on these views, the change can cut the copy overhead that PyTorch says previously moved more data than the reduction needed to read.

How the reduction path changes

The plan uses at::collapse_dims to express an input as rows and a run, then has reduction_flat_strided traverse that view with a fixed stride. Threadgroups grid-stride across rows, while lanes process elements in each run. Short rows can be packed into a threadgroup; long runs can be divided into equal chunks to distribute a strided one-dimensional view across the grid.

Operations and kernel sizing

PyTorch wrote the path against its shared operation template, covering sum, mean, nansum, count_nonzero, max, min, all, any, var and std. The strided kernel has a 2,048-element threadgroup budget, versus 8,192 for Flat: PyTorch says short-row strided work is latency-bound, and a 16K-element view using the Flat budget ran on two threadgroups and lost to the copy path.

Reported M4 Pro results

Against the existing .contiguous()-plus-Flat path, PyTorch reported fp16 speedups of 1.92x for sum and 1.91x for max on [4096, 8192][:, ::2], reducing the respective times from 534 to 278 microseconds and 536 to 280 microseconds. On [8192, 4096][::2], fp16 sum fell from 407 to 151 microseconds (2.69x), while max fell from 403 to 149 microseconds (2.71x).

The largest listed result was fp32 max on [8192, 4096][::2], from 934 to 303 microseconds, or 3.09x; fp32 sum on the same input went from 932 to 303 microseconds, or 3.07x. Gains were not universal: fp16 sum on [1024, 2048][::2] measured 18.2 microseconds versus 17.3 microseconds for the prior path, a 0.95x result. PyTorch says x[::2] now costs the same as an equivalent dense reduction, while x[:, ::2] is about twice as expensive as dense due to stride-2 cache-line access. Dense inputs are unaffected.

Fallbacks and validation

Layouts that collapse to three or more blocks, and runs that cannot be split evenly, including prime lengths, retain the contiguous-copy fallback. PyTorch reported 1,515 passing tests and 113 expected failures in its MPS test selection, plus a 540-case comparison against CPU spanning 12 operations, six dtypes and layouts including inner and outer slices, transposed and offset views, stride-0 expansion, and fallback cases.

Source: PyTorch release notes.

Definition. FlatStrided is a PyTorch MPS execution plan for supported full reductions on eligible strided tensor views.

BenchmarkReported result
fp16 sum, [4096, 8192][:, ::2]534 to 278 microseconds (1.92x)
fp16 max, [4096, 8192][:, ::2]536 to 280 microseconds (1.91x)
fp16 sum, [8192, 4096][::2]407 to 151 microseconds (2.69x)
fp16 max, [8192, 4096][::2]403 to 149 microseconds (2.71x)
fp32 max, [8192, 4096][::2]934 to 303 microseconds (3.09x)
fp16 sum, [1024, 2048][::2]18.2 versus 17.3 microseconds (0.95x)

Key takeaways

  • Eligible non-contiguous MPS tensor views can now be reduced in place instead of materializing a contiguous copy.
  • The path supports sum, mean, nansum, count_nonzero, max, min, all, any, var and std.
  • Reported M4 Pro results included 1.92x fp16 sum and 1.91x fp16 max speedups on [4096, 8192][:, ::2].
  • The largest listed result was fp32 max on [8192, 4096][::2], improving from 934 to 303 microseconds.
  • Dense inputs are unaffected, while some layouts and unevenly splittable runs retain the contiguous-copy fallback.

FAQ

What does PyTorch’s FlatStrided MPS path change?

It allows eligible full reductions on non-contiguous MPS tensors to traverse strided views directly rather than creating a contiguous copy first.

Which reductions does the strided path cover?

It covers sum, mean, nansum, count_nonzero, max, min, all, any, var and std.

Are all strided layouts supported?

No. Layouts that collapse to three or more blocks and runs that cannot be split evenly, including prime lengths, keep the contiguous-copy fallback.

How did the new path perform on M4 Pro?

PyTorch reported speedups up to 3.09x in the listed comparisons, though one fp16 sum case measured 0.95x versus the prior path.