How exactly the map-reduce will happen though? Won’t you need to do it host-side, or make a lot of reads and writes?
Also, doesn’t it mean that you forgo batching?
Map-reduce is implemented as a rolling calc, see: online softmax in FlashAttention kernels.
Map-reduce is implemented as a rolling calc, see: online softmax in FlashAttention kernels.