DataWhale组队学习笔记--llm-algo-leetcode(五)

文章目录

  • [vLLM PagedAttention(vLLM 分页注意力)](#vLLM PagedAttention(vLLM 分页注意力))
    • [1. 痛点:传统 KV Cache 的内存刺客](#1. 痛点:传统 KV Cache 的内存刺客)
    • [2. 破局思路:引入操作系统的虚拟内存](#2. 破局思路:引入操作系统的虚拟内存)
    • [3. PagedAttention 的工作机制](#3. PagedAttention 的工作机制)
    • [4. PagedAttention 带来的核心红利](#4. PagedAttention 带来的核心红利)
  • [SGLang RadixAttention(SGLang基数注意力)](#SGLang RadixAttention(SGLang基数注意力))
    • [1. 核心痛点:重复 Prefill 与跨请求浪费](#1. 核心痛点:重复 Prefill 与跨请求浪费)
    • [2. 核心原理:用基数树(Radix Tree)自动管理 KV Cache](#2. 核心原理:用基数树(Radix Tree)自动管理 KV Cache)
    • [3. 典型应用场景](#3. 典型应用场景)
    • [4. 对比:PagedAttention vs. RadixAttention](#4. 对比:PagedAttention vs. RadixAttention)

vLLM PagedAttention(vLLM 分页注意力)

在大语言模型(LLM)的部署和推理优化中,vLLM 提出的 PagedAttention(分页注意力机制) 是一个里程碑式的突破。它直接解决了大模型在生成长文本时,显存(GPU Memory)被严重浪费的核心痛点。

以下是 PagedAttention 的核心原理及其实战机制的拆解:

1. 痛点:传统 KV Cache 的内存刺客

在自回归(Autoregressive)生成过程中,大模型需要保存之前生成的每个 Token 的 Key 和 Value 张量,以避免每次生成新词时重复计算。这些被缓存的数据就叫 KV Cache

传统推理框架在管理 KV Cache 时面临严重的内存浪费问题:

  • 预分配导致内部碎片:系统通常会为了以防万一,预先为每个请求分配一段连续的、达到模型最大允许长度的显存空间(例如 2048 或 4096 tokens)。但很多请求实际生成的长度远达不到最大值,多出来的显存就被白白占用了。
  • 无法高效共享:在并行采样(如 Beam Search)中,同一个 Prompt 会产生多个不同的输出分支。传统的连续内存分配无法让这些分支共享 Prompt 阶段的 KV Cache,导致一份前缀被复制了多份。

据统计,传统的连续内存管理机制会导致高达 60% - 80% 的 KV Cache 显存被浪费。

2. 破局思路:引入操作系统的虚拟内存

vLLM 团队从操作系统的虚拟内存分页机制(Virtual Memory Paging)中汲取了灵感。

在操作系统中,程序认为自己拥有一块连续的"逻辑内存",但实际上这些内存在物理上是被切分成一个个固定大小的"页(Pages)",并可以分散存储在物理内存的各个角落。系统通过"页表(Page Table)"来记录逻辑地址到物理地址的映射关系。

3. PagedAttention 的工作机制

PagedAttention 将这种分页思想直接移植到了 LLM 的注意力机制计算中:

  1. KV Blocks(KV 块):系统不再为每个请求分配一整段连续的 KV 显存,而是将其划分为一个个固定大小的块(Block)。每个 Block 包含固定数量的 Token 的 KV 向量(例如,每个 Block 只装 16 个 Token)。
  2. 非连续物理存储 :这些 Block 在 GPU 的显存中不需要连续存放,哪里有空闲空间就分配在哪里。
  3. Block Table(块表):vLLM 维护着一个中心化的块表,负责将每个请求连续的"逻辑 KV 块"映射到显存中散落的"物理 KV 块"。

当模型进行 Attention 计算时,PagedAttention 算法会在底层自动查找块表,跨越这些非连续的物理块,准确地提取出需要的 KV 向量完成矩阵乘法。这个过程对上层模型是完全透明的。

4. PagedAttention 带来的核心红利

  • 几乎消除显存浪费:因为是按需动态分配(用完一个 Block 再分配下一个),显存的内部浪费被严格控制在最后一个没有装满的 Block 内(通常浪费率低于 4%)。
  • 吞吐量(Throughput)翻倍 :节省下来的海量显存,可以用来同时容纳更多并发请求的 KV Cache,从而显著增大 Batch Size。在相同的硬件下,vLLM 的吞吐量通常能达到传统框架的 2 - 4 倍
  • 写时复制(Copy-on-Write)实现内存共享:对于多个请求共享同一个 Prompt,或者执行 Beam Search 这种有共同前缀的场景,不同请求在 Block Table 中只需指向相同的物理块。只有当它们各自生成不同的后续 Token 时,系统才会分配新的物理块(写时复制)。这使得复杂推理场景下的显存占用呈指数级下降。

代码实现:

python 复制代码
import torch
from typing import List
python 复制代码
class Request:
    def __init__(self, request_id: int, prompt_len: int):
        self.request_id = request_id
        self.seq_len = prompt_len
        self.block_table: List[int] = []

class KVCacheManager:
    def __init__(self, num_blocks: int, block_size: int, head_dim: int):
        self.num_blocks = num_blocks
        self.block_size = block_size
        self.head_dim = head_dim
        
        # TODO 1: 模拟预分配一块大显存池
        self.physical_kv_cache = torch.zeros(num_blocks, block_size, head_dim)
        
        # 跟踪哪些物理块被占用了
        self.free_blocks: List[int] = list(range(num_blocks))

    def allocate_for_prefill(self, req: Request):
        """
        请求刚进来时 (Prefill阶段),为它的 Prompt 长度分配所需的全部 Block
        """
        # TODO 2: 计算需要的 block 数量(向上取整)
        needed_blocks = (req.seq_len + self.block_size - 1) // self.block_size
        
        # TODO 3: 从 free_blocks 中弹出对应数量的 block 索引
        if len(self.free_blocks) < needed_blocks:
            raise RuntimeError("OOM")
        
        for _ in range(needed_blocks):
            block_id = self.free_blocks.pop(0)
            req.block_table.append(block_id)

    def allocate_for_decode(self, req: Request):
        """
        自回归生成时 (Decode阶段),检查序列长度。
        如果当前最后一个 Block 满了,则按需分配 1 个新 Block。
        """
        req.seq_len += 1
        
        # TODO 4: 判断是否需要新的 Block
        is_new_block_needed = (req.seq_len % self.block_size) == 1
        
        if is_new_block_needed:
            if not self.free_blocks:
                raise RuntimeError("OOM")
            block_id = self.free_blocks.pop(0)
            req.block_table.append(block_id)

    def get_physical_cache(self, req: Request) -> torch.Tensor:
        """
        根据块表,把不连续的物理块"拼凑"成逻辑上连续的 KV Cache
        """
        # TODO 5: 根据 req.block_table 的索引,从物理池中提取对应的块
        blocks = [self.physical_kv_cache[block_id] for block_id in req.block_table]
        cat_blocks = torch.cat(blocks, dim=0)
        
        # 只截取真实 seq_len 长度返回
        return cat_blocks[:req.seq_len]
python 复制代码
# 运行此单元格以测试你的实现
def test_paged_attention_manager():
    try:
        # Case 1: 典型 Prefill + Decode + Cache 拼装
        manager = KVCacheManager(num_blocks=10, block_size=4, head_dim=64)
        print("初始化内存池...")

        req1 = Request(request_id=1, prompt_len=6)
        manager.allocate_for_prefill(req1)
        assert len(req1.block_table) == 2, "长度 6 的请求应分配 2 个 Block!"
        assert len(manager.free_blocks) == 8, "池中应该剩下 8 个空闲块!"
        print(f"✅ Prefill 测试通过!Req1 分配的块表: {req1.block_table}")

        manager.allocate_for_decode(req1)
        assert len(req1.block_table) == 2, "生成第 7 个 token 时不应该分配新块!"

        manager.allocate_for_decode(req1)
        manager.allocate_for_decode(req1)
        assert len(req1.block_table) == 3, "生成第 9 个 token 时应当分配了第 3 个新块!"
        assert len(manager.free_blocks) == 7, "池中应该剩下 7 个空闲块!"
        print(f"✅ Decode 动态分配测试通过!Req1 最新块表: {req1.block_table}")

        for block_id, value in zip(req1.block_table, [1.0, 2.0, 3.0]):
            manager.physical_kv_cache[block_id].fill_(value)
        cache = manager.get_physical_cache(req1)
        assert cache.shape == (9, 64), f"拼装出来的连续 Cache 形状不对,应为 (9, 64),实为 {cache.shape}"
        assert torch.all(cache[:4] == 1.0), "第 1 个 Block 未正确拼装!"
        assert torch.all(cache[4:8] == 2.0), "第 2 个 Block 未正确拼装!"
        assert torch.all(cache[8:] == 3.0), "第 3 个 Block 的截断拼装不正确!"
        print("✅ Cache 拼装测试通过!多块物理缓存被正确恢复为逻辑连续序列。")

        # Case 2: 恰好跨越 block 边界时,Decode 应该分配新块,并正确截断最后一块
        manager2 = KVCacheManager(num_blocks=4, block_size=4, head_dim=8)
        req2 = Request(request_id=2, prompt_len=4)
        manager2.allocate_for_prefill(req2)
        assert len(req2.block_table) == 1, "长度 4 的请求应只分配 1 个 Block!"
        manager2.allocate_for_decode(req2)
        assert len(req2.block_table) == 2, "长度 5 的请求应分配第 2 个 Block!"
        manager2.physical_kv_cache[req2.block_table[0]].fill_(7.0)
        manager2.physical_kv_cache[req2.block_table[1]].fill_(8.0)
        cache2 = manager2.get_physical_cache(req2)
        assert cache2.shape == (5, 8), f"拼装出来的连续 Cache 形状不对,应为 (5, 8),实为 {cache2.shape}"
        assert torch.all(cache2[:4] == 7.0), "边界块的前 4 个 token 不正确!"
        assert torch.all(cache2[4:] == 8.0), "边界块的最后 1 个 token 不正确!"
        print("✅ 边界分配与截断测试通过!")

        # Case 3: OOM 分支必须抛出 RuntimeError
        oom_manager = KVCacheManager(num_blocks=1, block_size=4, head_dim=8)
        oom_req = Request(request_id=3, prompt_len=5)
        try:
            oom_manager.allocate_for_prefill(oom_req)
        except RuntimeError as e:
            assert "OOM" in str(e), "OOM 异常信息不正确!"
            print("✅ OOM 测试通过!")
        else:
            raise AssertionError('显存池不足时应该抛出 RuntimeError("OOM")!')

        print("\n✅ All Tests Passed! PagedAttention 内存管理逻辑验证通过。")

    except NotImplementedError:
        print("请先完成 TODO 部分的代码!")
        raise
    except (AttributeError, NameError, TypeError, ValueError, AssertionError, RuntimeError) as e:
        if isinstance(e, AttributeError):
            print("代码未完成,无法找到必要的属性")
        elif isinstance(e, NameError):
            print("代码可能未完成,导致变量为 NoneType。")
        elif isinstance(e, TypeError):
            print("代码可能未完成,导致变量为 NoneType。")
        elif isinstance(e, ValueError):
            print("代码可能未完成,导致了张量维度错误")
        elif isinstance(e, AssertionError):
            print("代码可能未完成,导致了断言失败")
        elif isinstance(e, RuntimeError):
            print("代码可能未完成,导致了运行时错误")
        else:
            print("代码可能未完成,导致了断言失败")
        raise NotImplementedError("请先完成 TODO 部分的代码!") from e
    except Exception as e:
        print(f"❌ 测试失败: {e}")
        raise


test_paged_attention_manager()

结果:

SGLang RadixAttention(SGLang基数注意力)

如果说 vLLM 的 PagedAttention 解决了"单请求内部(Intra-request)显存碎片化"的问题,那么 SGLang 提出的 RadixAttention(基数注意力机制) 则更进一步,解决了"跨请求/多轮交互间(Inter-request)KV Cache 的自动复用与生命周期管理"的难题。


1. 核心痛点:重复 Prefill 与跨请求浪费

在大模型实际应用场景中,大量请求都包含重叠的前缀(Overlapping Prefixes)

  • 多轮对话(Multi-turn Chat):第 2 轮对话的输入包含第 1 轮的 Prompt 和 Answer。
  • Agent / 复杂工作流:在 Tree-of-Thought(思维树)或 Monte Carlo 搜索中,多个分支共享相同的推导历史。
  • 固定 System Prompt / Few-shot 示例:万级请求共享同一段很长的规则说明或示例。

在传统推理引擎中,哪怕请求之间有 90% 的 Token 完全相同,新请求到来时系统依然要对其前缀重新做一次 Prefill(预填充计算),不仅浪费算力,还会导致首包延迟(TTFT, Time-To-First-Token)居高不下。


2. 核心原理:用基数树(Radix Tree)自动管理 KV Cache

SGLang 没有采用复杂的全局 Hash 表,而是引入了计算机科学中经典的 Radix Tree(基数树/压缩前缀树) 来作为 KV Cache 的索引结构。

在 RadixTree 中:

  1. 边(Edges)与节点(Nodes) :保存连续的 Token 序列(如 "You are a helpful assistant...")。
  2. 指针与物理映射:节点直接关联底层物理显存中的 KV Cache 块(通常结合了类似 PagedAttention 的 Block 机制)。
  3. 动态生命周期
  • 自动匹配(Prefix Matching) :新 Prompt 进来时,在树中从根节点向下做最长前缀匹配。匹配到的部分直接复用 KV Cache,彻底跳过这部分的 Prefill 计算,模型只需要对"新后缀"进行计算。
  • 动态分裂(Split & Insert):当新请求在某个节点中途出现分叉时,原有节点会自动拆分为一个公共父节点和两个子节点。
  • LRU 淘汰(LRU Eviction):当 GPU 显存满载时,系统会根据 LRU(最近最少使用)策略,优先删除树的叶子节点(Leaf Nodes)对应的 KV Cache,并回收显存,直到空间足够。

3. 典型应用场景

RadixAttention 的最大优势在于 "零配置全自动" ------ 开发者不需要手动维护复杂的 Cache 清单,系统会在后台自动识别模式并完成 KV 共享。

在以下四种常见模式中,RadixAttention 能带来数倍的吞吐与延迟优化:

  1. Few-shot Learning(少样本提示):多个请求共享相同的示例 Prompt,前缀 Cache 命中率接近 100%。
  2. Multi-turn Chat(多轮交互):随对话轮数增加,历史上下文全部在树上,后续轮次只需 Prefill 用户刚发送的一句话。
  3. Self-consistency(采样一致性):单 Prompt 产生多个采样分支,前缀只计算一次。
  4. Tree-of-Thought(思维树搜索):多路探索任务中,所有子分支自动共享根节点与父节点的搜索历史。

4. 对比:PagedAttention vs. RadixAttention

维度 PagedAttention (vLLM) RadixAttention (SGLang)
解决的核心问题 解决单请求内物理显存离散化与碎片问题 解决跨请求间自动前缀复用与调度问题
索引数据结构 扁平的物理块映射表(Block Table) 动态层级基数树(Radix Tree)
前缀匹配机制 主要是单请求/显式配置的前缀缓存 全自动最长前缀匹配(Automatic Prefix Caching)
缓存回收策略 请求结束即立即释放(或简单保留) 基于树结构的 LRU 延迟释放(作为全局 Cache 池)
性能优势点 极大提高 Batch Size 与 GPU 利用率 极大地降低多轮/Agent 场景的 TTFT (首字延迟)

一句话总结:PagedAttention 提供了高效的"物理内存物理块"管理,而 RadixAttention 在其之上盖了一层"逻辑前缀索引树",两者结合成为了现代大模型推理引擎(如 SGLang、vLLM v1/v2 架构)的标准配置。

代码实现:

python 复制代码
import torch
python 复制代码
class TreeNode:
    def __init__(self,key_tokens):
        self.key_tokens = key_tokens #这条边上的Token序列(如[101,532,789])
        self.children = []  #子节点列表
        self.kv_cache_ptr = None  #模拟指向物理KV Cache的指针
    
class SimpleRadixCache:
    def __init__(self):
        #根节点是空的
        self.root = TreeNode([])

    def insert(self,tokens):
        node = TreeNode(tokens)
        self.root.children.append(node)

    def _lcp_len(self,cached_tokens,prompt_tokens):
        match_len = 0
        #TODO1:逐个token计算最长公共前缀长度,遇到不相等时立刻停止
    
        match_len = 0
        while match_len < len(cached_tokens) and match_len < len(prompt_tokens):
            if cached_tokens[match_len] == prompt_tokens[match_len]:
                match_len += 1
            else:
                break
        return match_len

    def match_prefix(self,prompt_tokens):
        best_match_len = 0

        #TODO2:遍历self.root.children,更新最长匹配前缀长度
        for child in self.root.children:
            match_len = self._lcp_len(child.key_tokens,prompt_tokens)
            if match_len > best_match_len:
                best_match_len = match_len
        return best_match_len

    def split_prompt(self,prompt_tokens):
        #TODO3:先找命中长度,再拆出前缀和后缀
        hit_len = self.match_prefix(prompt_tokens)
        hit_prefix = prompt_tokens[:hit_len]
        miss_suffix = prompt_tokens[hit_len:]

        return hit_prefix,miss_suffix,hit_len
python 复制代码
# 测试你的实现
def test_radix_attention():
    try:
        cache = SimpleRadixCache()
        cache.insert([0, 1, 2, 3])
        cache.insert([0, 1, 2, 3, 4])
        cache.insert([9, 9, 9])

        # 1. 基础 LCP 检查
        assert cache._lcp_len([1, 2, 3], [1, 2, 4]) == 2, "LCP 计算失败!"
        assert cache._lcp_len([7, 8], [7, 8, 9, 10]) == 2, "完整前缀匹配失败!"
        print("✅ 最长公共前缀计算正确!")

        # 2. 多候选路径下,应该选择最长命中前缀
        match_len = cache.match_prefix([0, 1, 2, 3, 4, 5])
        assert match_len == 5, "匹配失败!应该命中最长的 5 个 token 前缀。"
        assert cache.match_prefix([7, 6, 5]) == 0, "错误匹配!不该匹配到任何东西。"
        print("✅ 多路径前缀命中选择正确!")

        # 3. 前缀拆分验证
        hit_prefix, miss_suffix, hit_len = cache.split_prompt([0, 1, 2, 3, 4, 5])
        assert hit_len == 5, "Hit Length 计算错误!"
        assert hit_prefix == [0, 1, 2, 3, 4], "可复用前缀拆分错误!"
        assert miss_suffix == [5], "待重算后缀拆分错误!"

        hit_prefix2, miss_suffix2, hit_len2 = cache.split_prompt([7, 6, 5])
        assert hit_len2 == 0, "无命中时 Hit Length 应为 0!"
        assert hit_prefix2 == [], "无命中时前缀应为空!"
        assert miss_suffix2 == [7, 6, 5], "无命中时后缀应保持原样!"
        print("✅ 前缀拆分与回退逻辑正确!")

        print("\n 所有测试通过!这正是 SGLang 让大模型推理首字响应飞升 10 倍的底层秘密!")

    except NotImplementedError:
        print("请先完成 TODO 部分的代码!")
        raise
    except (AttributeError, NameError, TypeError, ValueError, AssertionError, RuntimeError) as e:
        if isinstance(e, AttributeError):
            print("代码未完成,无法找到必要的属性")
        elif isinstance(e, NameError):
            print("代码可能未完成,导致了变量未定义")
        elif isinstance(e, TypeError):
            print("代码可能未完成,导致了操作错误")
        elif isinstance(e, ValueError):
            print("代码可能未完成,导致了张量维度错误")
        elif isinstance(e, AssertionError):
            print("代码可能未完成,导致了断言失败")
        elif isinstance(e, RuntimeError):
            print("代码可能未完成,导致了运行时错误")
        else:
            print("代码可能未完成,导致了断言失败")
        raise NotImplementedError("请先完成 TODO 部分的代码!") from e
    except Exception as e:
        print(f"❌ 发生未知异常: {e}")
        raise


test_radix_attention()

结果:

相关推荐
是上好佳佳佳呀2 小时前
【机器学习|DAY03】K近邻算法(KNN)笔记
笔记·机器学习·近邻算法
To_OC9 小时前
大模型蒸馏是啥?说白了就是大厨带徒弟的学问
人工智能·llm·agent
陳陈陳9 小时前
从“胡说八道”到“妙笔生花”:我用LangChain手搓了一个可控AI写作流(Temperature+TopK调参指南)
langchain·llm
世人万千丶9 小时前
鸿蒙Flutter Flex多子组件权重分配
学习·flutter·华为·harmonyos·鸿蒙
冬奇Lab11 小时前
AI 评测系列(03):LLM-as-Judge——让 LLM 评价 LLM 的正确姿势
人工智能·llm
圣光SG12 小时前
Servlet学习笔记
笔记·学习·servlet
奋发向前wcx12 小时前
y1,y2总复习笔记5 2026.7.19
数据结构·笔记·算法
@Mike@14 小时前
02-数据库学习笔记(SQL引擎)
数据库·笔记·学习
Darling噜啦啦14 小时前
揭秘 LLM 的随机性黑盒:从 Temperature + Top-K 到 LangChain Chain 工作流实战
langchain·llm