Torch-PIM speeds up deep learning by smartly moving work to memory cores

Torch-PIM: Automated Profile-Guided PIM Offloading for PyTorch

Hardware ArchitectureDistributed, Parallel, and Cluster Computing

Summary

Deep learning often slows down because data has to move a lot between memory and processors. The authors created Torch-PIM, a tool that automatically decides which parts of a deep learning program should run close to memory to save time. It looks carefully at the program’s loops and uses information from running the program to make these choices. Using Torch-PIM can make common deep learning tasks run several times faster compared to only using the main processor.

What this means in practice

  • For deep learning engineers: Automatically decide when to run parts of PyTorch models on processing-in-memory hardware for faster training and inference workloads.
  • For data center architects: Design systems that deploy PyTorch with Torch-PIM to reduce data movement bottlenecks in large-scale AI workloads.

Authors

Heeeon Lee, Hyunwoo Nam, Junyong Heo, Hyunmo Sung, Jay Hwan Lee, Yeonsoo Kim, Seongho Jeong, Shinhyung Yang, Bernd Burgstaller

Abstract

Modern deep learning (DL) workloads are limited by data movement, and processing-in-memory (PIM) targets this bottleneck by placing compute units near the memory. However, PyTorch and other DL frameworks lack compiler support for making this decision on the code they lower: existing offloading frameworks target hand-written C/C++ programs, while those that address DL fix the candidate set to a list of operator types before lowering. We present Torch-PIM, a compiler framework that uses profile-guided optimization (PGO) to decide host-versus-PIM placement over the loop nests that progressive lowering materializes. Every parallel loop nest the pipeline emits enters the candidate space, and each is assessed in two stages: the amount of work it carries, and its memory boundedness. Every quantity the assessment consumes is profiled on the host or obtained from the multi-level intermediate representation (MLIR) of the code. Across PIM configurations of 32 to 128 cores, Torch-PIM's offloading decisions yield speedups of up to 8.6x on tensor operators, 2.9x on MLP, 4.4x on Attention, 5.1x on GPT-J-6B, and 3.6x on LLaMA-7B over CPU-only execution.