尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

DFlash源码导读(二):Qwen3DFlashAttention注意力实现全解析

DFlash源码导读(二):Qwen3DFlashAttention注意力实现全解析 DFlash源码导读二Qwen3DFlashAttention注意力实现全解析【免费下载链接】dflashDFlash: Block Diffusion for Flash Speculative Decoding项目地址: https://gitcode.com/GitHub_Trending/df/dflashDFlash是一款专为大模型推理加速设计的轻量级block diffusion块扩散草稿模型用于推测解码主模型target model负责逐 token 校验DFlash 则一次性并行预测出整块 token从而显著提升生成速度。本文带你逐行读懂 DFlash 最核心的Qwen3DFlashAttention注意力实现——双流 Key/Value 拼接、非因果块注意力、RoPE 位置对齐与注意力后端切换四大设计。 DFlash 项目结构注意力源码在哪里文件作用dflash/model.pyPyTorch 版草稿模型本文主角Qwen3DFlashAttention在此dflash/model_mlx.pyMLX 版草稿模型对应DFlashAttention注意力实现dflash/benchmark.py性能基准脚本gsm8k、math500、humaneval 等数据集README.mdvLLM / SGLang / Transformers / MLX 四种部署方式的快速上手Qwen3DFlashAttention完整定义在 dflash/model.py#L185-L255只有约 70 行却浓缩了 4 个关键设计双流 Key/Value上下文来自目标模型噪声来自块内 token 本身非因果块注意力块内 token 互相可见支持整块并行预测全窗口 RoPE一段 cos/sin 同时覆盖上下文与块位置自动对齐弹性注意力后端一行代码在 eager / SDPA / FlashAttention 之间切换 看懂注意力数据流context 与 noise 从哪里来在动手读代码前先建立整体心智模型。DFlash 草稿模型每轮只处理两个输入流target_hidden上下文 ── k_proj/v_proj ──▶ K_ctx / V_ctx noise_embedding噪声 ── q/k/v_proj ──▶ Q / K_noise / V_noise │ K cat(K_ctx, K_noise) V cat(V_ctx, V_noise) │ 写入草稿 KV cache → 非因果 attention(Q, K, V) → o_proj上下文流目标模型中间层的 hidden states经 extract_context_feature 抽取按 build_target_layer_ids 的规则选层单层草稿取目标模型正中间一层多层则在第 1 层到倒数第 3 层之间均匀采样再经 DFlashDraftModel.forward 中的fc线性投影 RMSNorm 得到target_hidden。草稿模型不重新嵌入历史文本直接借用目标模型的特征省掉一次完整编码。噪声流当前块的 token1 个已知锚点 若干 mask token用目标模型的 embedding 表嵌入见 dflash_generate 中target.model.embed_tokens(block_output_ids)。 逐段拆解 forward五步读懂 Qwen3DFlashAttention1️⃣ QKV 投影与 QK-RMSNormself.q_norm Qwen3RMSNorm(self.head_dim, epsconfig.rms_norm_eps) self.k_norm Qwen3RMSNorm(self.head_dim, epsconfig.rms_norm_eps)dflash/model.py#L207-L208Q、K 在进入注意力前都在head_dim维做 RMSNorm这是 Qwen3 的 QK-Norm 设计保证注意力分数稳定head_dim优先取 config 中的显式值否则回退为hidden_size // num_attention_headsL190K/V 头数少于 Q 头数即 GQA 分组查询注意力num_key_value_groups num_attention_heads // num_key_value_headsL191省显存又省带宽。2️⃣ 双流 Key/Value 拼接DFlash 最独特的地方k_ctx self.k_proj(target_hidden) # 上下文 key来自目标模型特征 k_noise self.k_proj(hidden_states) # 噪声 key来自块内 token 自身 v_ctx self.v_proj(target_hidden) v_noise self.v_proj(hidden_states) k torch.cat([k_ctx, k_noise], dim1) # 拼成一条长序列 v torch.cat([v_ctx, v_noise], dim1)dflash/model.py#L226-L231拼接后 K/V 长度为ctx_len q_len即全部已接受上下文 当前块。上下文 KV 会在下一轮被直接写入草稿 KV cache 复用块内 KV 则是临时的后面会讲如何裁剪。3️⃣ RoPE 位置对齐一段 cos/sin 管两段位置cos, sin position_embeddings q, k apply_rotary_pos_emb(q, k, cos, sin)apply_rotary_pos_emb巧妙之处在于 apply_rotary_pos_emb外层传入的position_ids覆盖整个窗口长度 ctx_len q_len而 query 只截取最后 q_len 段的 cos/sin对应块的位置key 则使用整段上下文 key 落在其历史真实位置噪声 key 紧随其后。这样无需任何手工簿记query 与 key 的全局位置天然对齐目标模型。4️⃣ KV cache只保留被接受的前缀if past_key_values is not None: k, v past_key_values.update(k, v, self.layer_idx, cache_kwargs)dflash/model.py#L236-L238本轮的上下文 噪声 KV 一次性写入草稿的DynamicCache。一轮验证结束后主循环调用past_key_values_draft.crop(start)L120把未通过验证的噪声 KV 和被拒 token 的 KV 一起裁掉——缓存长度始终等于已接受前缀长度内存零泄漏。5️⃣ 注意力后端一行切换 滑动窗口attn_fn eager_attention_forward if self.config._attn_implementation ! eager: attn_fn ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]dflash/model.py#L239-L241复用 transformers 的统一注意力注册表eager矩阵实现便于调试、sdpa、flash_attention_2等只改配置即可切换模型代码零改动sliding_window仅对layer_types为sliding_attention的层生效L209从而支持 Qwen3.5 这类全注意力 滑动窗口注意力混合的 SWA 草稿模型。⚡ 关键设计非因果的块扩散注意力self.is_causal FalseL194且调用时attention_mask通常为 None——块内每个 token 能看到全部上下文 块内其他所有 token这正是 block diffusion 的本质类比扩散模型的 denoising给定锚点 token mask token模型一次前向就把整块 token并行解码出来而不是逐个串行生成生成循环中 logits 取[:, 1 - block_size :, :]L112-L119丢弃第一个锚点位置其内容已知其余位置采样后直接填入block_output_ids[:, 1:]成为提交给目标模型验证的草稿块。对比传统自回归注意力只能向左看Qwen3DFlashAttention 则赋予块双向视野——以训练复杂度换取整块并行的速度收益。 MLX 版注意力实现对照DFlashAttention 是同一思想在 Apple SiliconMLX上的移植对比项PyTorchmodel.pyMLXmodel_mlx.py双流 KVtorch.cat拼接 K_ctx/K_noise同思路mx.concatenate滑动窗口sliding_window参数交给注意力内核预截断上下文至最近窗口 因果窗口 maskL85-L91位置编码cos 取最后 q_len 段显式传offsetcache.offset SL102-L104注意力内核transformers 统一注册表mx.fast.scaled_dot_product_attentionL115KV cacheDynamicCachecropKVCache/RotatingKVCache trim结论把 DFlash 移植到新框架时只需重新实现双流 KV 拼接 位置对齐 非因果注意力这一核心三件套其余映射到目标框架的注意力内核即可。 小结四个记忆点双流 K/V上下文来自目标模型中间层 hidden statesfc RMSNorm 投影噪声来自块内 token一次 cat 拼成长序列非因果块注意力is_causal False块内 token 互看实现块级并行预测——这就是 block diffusion全窗口 RoPE一段 cos/sin 覆盖上下文 块query 取尾段、key 用全段位置零簿记自动对齐工程友好注意力后端一行切换、支持滑动窗口混合层、草稿 KV cache 只保留已接受前缀。 延伸阅读草稿模型主类DFlashDraftModel推测解码主循环提案 验证 接受dflash_generateMLX 草稿模型与流式生成dflash/model_mlx.py性能基准脚本dflash/benchmark.py部署示例vLLM / SGLang / Transformers / MLXREADME.md【免费下载链接】dflashDFlash: Block Diffusion for Flash Speculative Decoding项目地址: https://gitcode.com/GitHub_Trending/df/dflash创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表