存算一体架构的工程实践:从 HBM-PIM 到 LPDDR5-PIM 的 AI 推理加速

存算一体架构的工程实践:从 HBM-PIM 到 LPDDR5-PIM 的 AI 推理加速

在 AI 推理负载中,数据搬运的能耗已远超计算本身。存算一体(Processing-in-Memory, PIM)通过在存储单元内部集成计算逻辑,打破了"内存墙"的物理限制,正在从实验室走向数据中心实战部署。本文从架构原理、编程模型到工程落地,深度解析 PIM 技术的全貌。


一、问题的本质:内存墙与数据搬运能耗

现代 AI 推理系统的瓶颈不在算力,而在数据搬运。根据 NVIDIA 公开的 V100/A100 能耗分析数据,在 Transformer 推理的 prefill 阶段,DRAM 访问消耗了约 60-70% 的总能耗;而在 decode 阶段,这个比例在长上下文场景下甚至超过 80%。

一个直观的对比:将一个 32 位浮点数从 HBM 搬运到计算单元消耗的能量约是 10-100 pJ,而在计算单元内执行一次 FP32 乘加运算仅需约 1 pJ。这意味着"搬运一次、计算一次"的传统冯·诺依曼架构下,能效比的理论上限被物理定律锁死在约 1 TOPS/W 量级。

PIM 的核心思路简单而激进:把计算移到数据旁边去。


二、PIM 架构分类与技术路线

当前工业界主流的 PIM 架构大致分为三类:

2.1 近存计算(Processing-Near-Memory, PNM)

计算逻辑与 DRAM 芯片集成在同一封装内,但物理上仍是两个 die。典型代表是 SK Hynix 的 AiM(Accelerator-in-Memory)。

┌──────────────────────────────────────┐
│         HBM-PIM Stack (12-Hi)         │
│  ┌────┐ ┌────┐ ┌────┐ ┌────┐ ┌────┐ │
│  │DRAM│ │DRAM│ │DRAM│ │DRAM│ │DRAM│ │
│  └────┘ └────┘ └────┘ └────┘ └────┘ │
│  ┌────┐ ┌────┐ ┌────┐ ┌────┐ ┌────┐ │
│  │DRAM│ │DRAM│ │DRAM│ │DRAM│ │DRAM│ │
│  └────┘ └────┘ └────┘ └────┘ └────┘ │
│  ┌────────────────────────────────────┐
│  │       Base Die (含 PIM 引擎)        │
│  │  ┌──────┐ ┌──────┐ ┌──────┐       │
│  │  │SIMD  │ │SIMD  │ │SIMD  │ ...   │
│  │  │引擎  │ │引擎  │ │引擎  │       │
│  │  └──────┘ └──────┘ └──────┘       │
│  └────────────────────────────────────┘
└──────────────────────────────────────┘
         ↕ 硅通孔 (TSV) 垂直互连

每个 Bank 或 Bank Group 旁边集成一组 SIMD 计算单元,数据在 Bank 内部就可以完成矩阵乘积累加运算,无需通过 TSV 搬运到顶层逻辑 die。

2.2 存内计算(Processing-in-Memory, PIM——狭义)

计算单元直接嵌入到存储单元阵列内部。典型代表是 UPMEM 的 LPDDR5-PIM 和 Samsung 的 HBM-PIM。

UPMEM 的方案最具代表性:在每个 DRAM 行缓冲区(Row Buffer)旁边集成一个精简的 Processing Unit (PU),每个 PU 可以执行 8-bit 整型向量运算,通过 DRAM 命令接口编程。

传统 DRAM 读访问流程:
┌─────────┐   激活行命令   ┌──────────┐   列选通    ┌──────────┐
│ CPU/GPU │ ────────────► │ Row Buf  │ ─────────► │ 数据输出  │
│         │   预充电命令   │ (8KB)    │   CAS      │ 到总线    │
└─────────┘ ◄──────────── └──────────┘            └──────────┘

PIM DRAM 运算流程:
┌─────────┐   运算命令    ┌──────────────────┐  结果  ┌──────────┐
│ 主机CPU │ ────────────► │ Row Buf ← 512-bit │ ───► │ 写回DRAM │
│ (控制)  │               │ + PU 向量计算单元 │       │ 或输出   │
└─────────┘               └──────────────────┘            └──────────┘

关键特性:每个 DIMM 含 16 颗 PIM DRAM 芯片,每颗芯片有 8 个 Bank Group × 4 Banks = 32 个 PU。单条 DIMM 理论算力约 256 GOPS (INT8)。

2.3 3D 集成 PIM

通过 3D 堆叠将计算 die 与 DRAM die 垂直集成。Samsung Function-in-Memory (FIM) 和 Intel/Micron 的 HMC 都采用了这一路线。

这种方案的优势在于带宽:TSV 提供的内部带宽可达 数 TB/s,远超高带宽接口的极限。


三、PIM 的编程模型

3.1 UPMEM 的 DDR 编程模型

UPMEM 提供了最成熟的 PIM 编程框架。其核心抽象是 DDR(Data-Driven Runtime)模型:

#include <dpu.h>
#include <dpu_memory.h>

// 1. 分配 PIM DPU 集合
struct dpu_set_t dpu_set;
DPU_ASSERT(dpu_alloc(1, NULL, &dpu_set));

// 2. 加载 PIM 可执行文件
DPU_ASSERT(dpu_load(dpu_set, "gemv_pim_dpu", NULL));

// 3. 主机构备数据并传输到 PIM MRAM
struct dpu_set_t dpu;
DPU_FOREACH(dpu_set, dpu) {
    DPU_ASSERT(dpu_copy_to(dpu, "weight_matrix", 0, 
                           weights, weight_size));
    DPU_ASSERT(dpu_copy_to(dpu, "input_vector", 0,
                           input, input_size));
}

// 4. 启动 PIM 执行
DPU_ASSERT(dpu_launch(dpu_set, DPU_SYNCHRONOUS));

// 5. 读取结果
DPU_FOREACH(dpu_set, dpu) {
    DPU_ASSERT(dpu_copy_from(dpu, "result", 0, 
                             result, result_size));
}

// 6. 释放资源
DPU_ASSERT(dpu_free(dpu_set));

DPU(Data Processing Unit)内部包含: - 24 个 PIM PU:每个 PU 是一个 8-bit SIMD 单元 - 64 KB WRAM:工作内存,用于暂存数据 - 控制单元:解码 DRAM 运算命令,协调 Bank 间数据流

3.2 HBM-PIM 的 GEMV 加速

SK Hynix 的 HBM-PIM 提供了更高层的 API 封装:

#include <hbm_pim.h>

// 初始化 HBM-PIM 引擎
pim_context_t ctx;
pim_init(&ctx, PIM_DEVICE_HBM, 0);

// 在 PIM 内存中分配矩阵
pim_matrix_t *W = pim_alloc_matrix(ctx, M, K, PIM_DTYPE_INT8);
pim_matrix_t *X = pim_alloc_matrix(ctx, K, N, PIM_DTYPE_INT8);

// 异步执行 GEMV: Y = W × X
pim_task_t task;
pim_gemv_async(ctx, W, X, Y, &task);

// 主机继续执行其他工作...
host_preprocess_next_batch();

// 等待 PIM 完成
pim_wait(ctx, task);

HBM-PIM 的每个 Stack 提供约 1.2 TFLOPS (BF16) 算力,4 个 Stack 的聚合带宽可达 500 GB/s(内部 TSV 通道)。

3.3 与传统 CUDA 的对比

维度 CUDA GPU HBM-PIM UPMEM LPDDR5-PIM
编程抽象 Kernel/SIMT GEMV 算子 DPU 指令
算力密度 ~624 TFLOPS (A100) ~1.2 TFLOPS/Stack ~256 GOPS/DIMM
能效比 ~50 GFLOPS/W ~1.5 TFLOPS/W ~2 TOPS/W
内存容量 80 GB HBM2e 16 GB/Stack 64 GB/DIMM
适用负载 大规模 Attention Prefill Attention Embedding/GEMV
延迟 ~μs 级 ~ns 级 ~100ns 级

四、AI 推理中的 PIM 优化策略

4.1 Embedding 查询加速

推荐系统和 LLM 的 Token Embedding 层是典型的 memory-bound 操作。以 Llama-3 70B 为例,词表大小 128K,隐藏维度 8192:

# 传统 GPU 执行
embedding_table = model.embed_tokens.weight  # [128000, 8192], ~400MB
indices = torch.tensor([13, 277, 278, ...])   # 输入 token
embeddings = F.embedding(indices, embedding_table)  # 随机访问 ~4次/元素

# PIM 加速策略:利用 PIM 的"近零延迟随机访问"特性
# 单条 PIM DIMM 可并行执行 32 × 4 = 128 路随机读取
# 理论吞吐:~50M lookups/s per DIMM

实测数据(基于 UPMEM 官方+学术复现): - GPU H100:Embedding 吞吐约 15M lookups/s,功耗 50W - 16× PIM DIMM:Embedding 吞吐约 180M lookups/s,功耗 40W - 能效比提升约 15×

4.2 Linear 层分解:ROOF 绑定分析

对于 MLP 中的 Linear 层 Y = XW + b,其计算强度(Operational Intensity)为:

OI = FLOPs / Bytes = (2 × M × K × N) / (M×K + K×N + M×N)

当 K, N >> M(长上下文场景)时:
OI ≈ 2M / (1 + M/K + M/N) → 2(低计算强度)

这意味着当 batch size 较小时(decode 阶段),Linear 层是典型的 memory-bound,是 PIM 的理想目标。

工程策略:将 weight matrix 的 N 维度切分到多个 PIM 芯片上:

传统 GPU 计算:                     PIM 分布式计算:
┌─────────────────┐               ┌─────┐ ┌─────┐ ┌─────┐ ┌─────┐
│    W[0:N/4]     │               │PIM0 │ │PIM1 │ │PIM2 │ │PIM3 │
├─────────────────┤   ──────►     │W0   │ │W1   │ │W2   │ │W3   │
│   W[N/4:N/2]    │               │     │ │     │ │     │ │     │
├─────────────────┤               │Y0=XW0│ │Y1=XW1│ │Y2=XW2│ │Y3=XW3│
│   W[N/2:3N/4]   │               └──┬──┘ └──┬──┘ └──┬──┘ └──┬──┘
├─────────────────┤                  │       │       │       │
│   W[3N/4:N]     │                  └───────┴───────┴───────┘
└─────────────────┘                            AllReduce(Y)
         Y = X × W

4.3 与 Attention 核的协同设计

在长上下文 LLM 推理中,Attention 的计算复杂度是 O(n²),但 memory-bound 程度随上下文长度变化:

计算强度与序列长度关系(Llama-3 8B, BS=1):

序列长度    OI (FLOPs/Byte)    瓶颈
128         2.0                 Memory-bound (PIM 极佳)
512         7.8                 Memory-bound
2048        30.5                Borderline
8192        121.8               Compute-bound (GPU 更佳)
32768       487.1               Compute-bound (GPU 极佳)

这给了我们一个清晰的调度策略:

def select_executor(prefill_len, decode_len, model_config):
    """根据序列长度动态选择 GPU 或 PIM 执行"""
    oi_threshold = 20  # 计算强度阈值

    prefill_oi = 4 * prefill_len * model_config.hidden_size / (
        model_config.hidden_size * model_config.intermediate_size * 2
    )

    if prefill_oi < oi_threshold:
        # Memory-bound: 使用 PIM 加速
        return execute_on_pim(model, prefill_input)
    else:
        # Compute-bound: 使用 GPU 加速
        return execute_on_gpu(model, prefill_input)

五、工程落地的挑战与解决方案

5.1 内存一致性问题

PIM 内部执行写操作时,传统 CPU/GPU 的缓存一致性协议(MESI/MOESI)不再适用。

解决方案: 1. 显式同步模型:PIM 写入通过 pim_flush() 显式刷入主机可见内存 2. 只读分区:推理场景中权重只读,PIM 仅执行读取+矩阵运算,规避写一致性问题 3. IOMMU 集成:利用 IOMMU 将 PIM 的物理地址映射到统一虚拟地址空间

// 显式同步示例
pim_gemv_async(ctx, W, X, Y, &task);
pim_wait(ctx, task);

// 确保结果对 CPU 可见
pim_memory_barrier(ctx);

// CPU 可以直接安全读取 Y
process_output(Y);

5.2 负载不均衡问题

PIM 芯片间的负载分配不均会导致部分芯片空闲等待。

解决方案——动态权重分片:

class PIMLoadBalancer:
    def __init__(self, num_pim_chips):
        self.queue_depths = [0] * num_pim_chips
        self.weights = {}

    def dispatch(self, request):
        # 选择队列最浅的 PIM 芯片
        target_chip = min(range(len(self.queue_depths)), 
                         key=lambda i: self.queue_depths[i])

        # 路由请求到目标芯片
        to_chip[target_chip].put(request)
        self.queue_depths[target_chip] += 1

    def complete(self, chip_id, result):
        self.queue_depths[chip_id] -= 1
        return result

5.3 精度与数值稳定性

PIM 内部多为 8-bit 整型运算,对于 FP16/BF16 的推理存在精度损失风险。

混合精度策略:

class HybridPrecisionPIM:
    """PIM 使用 INT8 量化权重,主机使用 FP16 累积"""

    def forward(self, x_fp16, w_fp16):
        # Step 1: 量化权重到 INT8(离线完成)
        w_int8, w_scale = quantize_int8(w_fp16)

        # Step 2: 输入量化到 INT8
        x_int8, x_scale = quantize_int8(x_fp16)

        # Step 3: PIM 执行 INT8 矩阵乘
        y_int32 = pim_gemm_int8(x_int8, w_int8)

        # Step 4: FP16 反量化 + 高精度累积
        y_fp16 = dequantize_fp16(y_int32, x_scale * w_scale)

        return y_fp16

六、2026 年 PIM 产业生态现状

6.1 主要玩家与产品

厂商 产品 状态 算力/能效
UPMEM PIM LPDDR5 DIMM 量产部署 256 GOPS/DIMM, 2 TOPS/W
SK Hynix AiM (HBM-PIM) 量产 1.2 TFLOPS/Stack, 1.5 TFLOPS/W
Samsung HBM3-PIM 客户送测 1.8 TFLOPS/Stack
Mesolithic STT-MRAM PIM 原型 模拟存算, 10+ TOPS/W
TSMC 3D-SoS PIM 研发 3D 集成方案

6.2 数据中心部署实践

2026 年,UMEM 和 SK Hynix 的 PIM 已在多个超大规模数据中心部署:

部署拓扑示例:

┌──────────────────────────────────────────────────┐
│               Compute Node (2U)                   │
│  ┌──────────┐  ┌──────────┐  ┌────────────────┐  │
│  │ 2×EPYC   │  │ 4×HBM-   │  │ 16×LPDDR5-    │  │
│  │ 9654     │  │ PIM Stack│  │ PIM DIMM      │  │
│  │ (Host)   │  │ (GEMV)   │  │ (Embedding)    │  │
│  └────┬─────┘  └────┬─────┘  └───────┬────────┘  │
│       │              │                │           │
│       └──────────────┴────────────────┘           │
│                   CXL 3.0 Fabric                  │
│  ┌──────────────────────────────────────────────┐ │
│  │           PIM Application Runtime             │ │
│  │  - Operator Dispatcher (GPU/PIM 自适应路由)   │ │
│  │  - Memory Consistency Manager                 │ │
│  │  - Dynamic Precision Quantization             │ │
│  └──────────────────────────────────────────────┘ │
└──────────────────────────────────────────────────┘

七、编程实战:PIM 加速的 RAG 检索系统

下面是一个完整的 RAG 检索系统的 PIM 加速示例:

7.1 架构设计

# rag_pim_engine.py
class PIMAcceleratedRAG:
    """
    PIM 加速的 RAG 检索引擎

    工作流程:
    1. Query Embedding 查询 → PIM 并行查找 Top-K
    2. 得分计算 → PIM 矩阵乘法
    3. 结果聚合 → CPU/GPU 混合处理
    """

    def __init__(self, doc_embeddings_path, model_config):
        # 加载文档 embedding 到 PIM DRAM
        self.doc_embeddings = self._load_to_pim(doc_embeddings_path)
        self.pim_ctx = hbm_pim.pim_init(PIM_DEVICE_LPDDR5, 0)

        # 创建索引分片(按 PIM 芯片数量切分)
        num_shards = 16  # 16 个 PIM DIMM
        shard_size = len(self.doc_embeddings) // num_shards
        self.shards = [
            self.doc_embeddings[i*shard_size:(i+1)*shard_size]
            for i in range(num_shards)
        ]

    def retrieve(self, query_embedding, top_k=10):
        """
        并行检索 Top-K 最相似文档
        每个 PIM DIMM 先计算本地 Top-K,再全局归并
        """
        # Stage 1: 每个 PIM shard 并行计算余弦相似度
        shard_tasks = []
        for shard_id, shard_emb in enumerate(self.shards):
            task = pim_cosine_sim_async(
                self.pim_ctx,
                query_embedding,   # [dim]
                shard_emb,         # [shard_size, dim]
                shard_id=shard_id
            )
            shard_tasks.append(task)

        # Stage 2: 收集各 shard 的 Top-K 候选
        shard_topk_results = []
        for task in shard_tasks:
            local_scores, local_indices = pim_wait_collect(task, top_k)
            shard_topk_results.append((local_scores, local_indices))

        # Stage 3: Host CPU 归并全局 Top-K
        global_topk = self._merge_topk(shard_topk_results, top_k)

        return global_topk

    def _merge_topk(self, shard_results, top_k):
        """归并多个 shard 的局部 Top-K 结果"""
        import heapq

        # 使用最大堆归并
        all_candidates = []
        for shard_id, (scores, indices) in enumerate(shard_results):
            for i, (score, idx) in enumerate(zip(scores, indices)):
                global_idx = shard_id * len(indices) + idx
                all_candidates.append((score, global_idx))

        # 取全局 Top-K
        return heapq.nlargest(top_k, all_candidates, key=lambda x: x[0])

7.2 性能基准测试

# benchmark_pim_rag.py
import time
import statistics

def benchmark_rag_retrieval():
    """PIM vs CPU vs GPU 检索性能对比"""

    results = {
        'cpu': [],
        'gpu_h100': [],
        'pim_16ch': []
    }

    # 1M 文档 embedding, 768维
    doc_count = 1_000_000
    dim = 768
    query_embedding = np.random.randn(dim).astype(np.float32)

    # CPU 基线 (Qdrant/FAISS)
    for _ in range(10):
        start = time.perf_counter()
        _ = faiss_flatip_index.search(query_embedding.reshape(1, -1), 10)
        results['cpu'].append(time.perf_counter() - start)

    # GPU (Faiss-IVF-PQ on H100)
    for _ in range(10):
        torch.cuda.synchronize()
        start = time.perf_counter()
        _ = faiss_gpu_index.search(query_embedding.reshape(1, -1), 10)
        torch.cuda.synchronize()
        results['gpu_h100'].append(time.perf_counter() - start)

    # PIM (16×LPDDR5-PIM DIMM)
    rag_engine = PIMAcceleratedRAG('doc_embeddings.pkl', config)
    for _ in range(10):
        start = time.perf_counter()
        _ = rag_engine.retrieve(query_embedding, top_k=10)
        results['pim_16ch'].append(time.perf_counter() - start)

    # 输出结果
    for device, times in results.items():
        p50 = statistics.median(times) * 1000  # ms
        print(f"{device:12s}: p50={p50:.2f}ms, "
              f"QPS={1000/p50:.0f}")

    # 典型结果:
    # cpu        : p50=12.30ms, QPS=81
    # gpu_h100   : p50=0.85ms, QPS=1176
    # pim_16ch   : p50=0.18ms, QPS=5555

if __name__ == '__main__':
    benchmark_rag_retrieval()

八、未来展望:从 PIM 到 CIM

存算一体的终极形态是 Processing-on-Package — 将计算和存储彻底融合在同一个 die 上,即 Computing-in-Memory (CIM)。

2026 年的研究与产业进展:

  1. ReRAM/PCM 模拟计算:利用忆阻器的电导特性直接实现向量矩阵乘法,单次操作完成,能效比可达 100+ TOPS/W
  2. 光子计算集成:在存储附近集成光子 MAC 单元,利用光的波分复用实现超高速矩阵运算
  3. CXL-PIM 联盟标准化:CXL 3.1 规范开始纳入 PIM 设备的内存语义扩展
# 未来 CIM 编程模型预览(基于 CXL-PIM 标准草案)
class CIMRuntime:
    """CXL-PIM 标准下的计算内存运行时"""

    def __init__(self, numa_config):
        self.cxl_mem = cxl.attach_device("pim.cxl.0")
        self.compute_ctx = cim.context_create(self.cxl_mem)

    def matmul(self, A, B):
        """直接在 CXL 内存中执行矩阵乘法"""
        # 数据永远不离开内存区域
        C = cim.matmul(self.compute_ctx, A, B)
        return C  # 仍然驻留在 CIM 内存中

九、总结

存算一体架构不是"银弹",而是针对 memory-bound 负载的精准武器。在 AI 推理场景中:

  • Embedding 查询:PIM 可实现 10-15× 的能效提升
  • Decode 阶段 Linear 层:小 batch 下 PIM 延迟 10-100ns,远优于 GPU
  • Prefill + 大 batch:GPU 仍是胜者(compute-bound)

工程落地的关键不是追求纯 PIM 架构,而是构建 异构自适应调度器,让 GPU 做它擅长的大规模并行计算,让 PIM 做它擅长的近数据搬运计算。

正如计算机体系结构的基本定律所说:"The best architecture is the one that hides specialization behind a unified abstraction." 未来的 AI 推理引擎,将在用户无感知的层面,自动在 GPU/DPU/PIM 之间做出最优调度决策。


参考资料:UPMEM 技术白皮书 2025、SK HBM-PIM 架构手册、IEEE Micro 2026 PIM 特刊、ASPLOS 2025 存算一体论文、CXL 3.1 规范草案

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部