Also, doesn’t it mean that you forgo batching?
How would you do map-reduce across multiple DIMMs w/o extra reads/writes?
PIM implies some sort of distributed compute, which can work for some cases, but I am not sure LLMs are one of them.
Each die-attached PIM accelerator computes online softmax for its own KVs. Then the central unit gathers the softmax intermediates, one intermediate per die, and uses those to compute the final softmax.
The PIM win is that we crater the memory traffic between the central accelerator and the memory dies for attention ops. Most of the attention bandwidth never leaves the memory.
This isn't "run the entire LLM in PIM", no - this is "offload the parts of LLM that benefit from PIM the most to PIM".