
TorchTune _export 目录全解析可被 torch.export 直接导出的注意力、KV Cache 与位置编码模块【免费下载链接】torchtunePyTorch native post-training library项目地址: https://gitcode.com/GitHub_Trending/to/torchtune在 TorchTune 中torchtune/modules/_export/目录提供了一组专门为模型导出export而改写的模块变体它们保持与参考模块完全一致的__init__()和forward()接口但内部规避了torch.export.export()无法处理的 Python 级数据依赖控制流与模块属性变更。读完本文你将理解这些导出友好模块与参考实现之间的逐项差异torch.cond分支改写、SDPA 子模块拆分、KV Cache 的clone()与cache_pos机制、如何一键替换模型中的 MHA 与视觉位置编码以及仓库中通过 eager / export / AOTI 三层测试保证数值等价性的验证方法。一、_export 目录的模块契约该目录的说明文档README定义了所有导出模块必须遵守的四条契约接口一致__init__()与forward()接受与torchtune中对应参考模块完全相同的参数输出一致返回值与参考模块相同除非 docstring 中另有说明导出兼容保证可以开箱即用地通过torch.export.export()导出AOT 兼容应当可以开箱即用地配合torch.aot_compile()使用。此外文档还规定了质量门槛与稳定性边界所有模块都必须被单元测试覆盖位于 tests/torchtune/modules/_export/这些测试会在每日构建与任何触及该目录的 PR 上运行这些模块处于活跃演进中subject to change so proceed with caution即接口不承诺长期稳定生产使用前应锁定版本并自行回归验证。目录当前包含的实现文件文件内容attention.py导出友好的MultiHeadAttention含 GQA/MQA 支持、SDPA子模块、批量替换函数kv_cache.py导出友好的KVCache继承自参考实现新增transpose_cache与clone()_position_embeddings.py导出友好的视觉位置编码TilePositionalEmbedding、TiledTokenPositionalEmbedding及替换函数install_requirements.sh环境安装脚本pip install torch2.6.0 torchvision0.21.0CPU nightly 索引版本要求从测试代码的版本门槛可以确认适用前提torch.cond相关用例标注unittest.skipUnless(torch_version_ge(2.6.0), reasontorch.cond only works for 2.6.0)export/AOTI 用例要求2.6.0.dev20241117之后的版本见 test_attention.py。install_requirements.sh也统一钉在torch2.6.0。因此使用该目录的模块至少需要 PyTorch 2.6.0。二、导出版 MultiHeadAttention与参考实现的差异attention.py 中的MultiHeadAttention改造自 torchtune/modules/attention.py 中的参考实现支持 MHA/GQA/MQAnum_kv_heads num_heads时为 GQA等于 1 时为 MQA。docstring 明确列出了三大改动将if y is None改写为torch.cond()条件判断改为y 的所有值是否都是 NaN以便torch.compile()/export 可以追踪由于torch.cond()的 true/false 两个分支都不能修改输入需要在torch.cond()之后把 kv 值再复制回 KV cache拆出独立的 SDPA 子模块SDPA 模块内部封装了 head 维度的 transpose 与 kv 头数的 expand便于导出程序的用户替换为自定义的高性能 SDPA 实现使用新的 KV cache把改为.add_以避免修改模块属性module attribute mutation并新增clone()方法。2.1 用 torch.cond 表达 KV cache 读写的分支参考实现中自回归解码时第二次调用forward会传yNone由 Python 的if y is None决定是计算新的 k/v 还是直接读缓存——这类根据运行期 Python 变量走不同图的控制流是 export 无法表达的核心障碍。导出版把该分支张量化约定y全为 NaN 表示不计算新 k/v从 KV cache 读取docstring 中明确写了 If all values are NaN, we read from kv cache实现attention.py 第 294-320 行def true_fn(y): kv_cache self.kv_cache.clone() return kv_cache.k_cache, kv_cache.v_cache, kv_cache.cache_pos def false_fn(y): k, v calculate_kv(y) kv_cache self.kv_cache.clone() kv_cache.update(k, v) return kv_cache.k_cache, kv_cache.v_cache, kv_cache.cache_pos k, v, cache_pos torch.cond( torch.isnan(y).all().item(), true_fn, false_fn, (y,) ) # 两个分支都不许改输入所以把结果再复制回真实 cache self.kv_cache.k_cache.copy_(k) self.kv_cache.v_cache.copy_(v) self.kv_cache.cache_pos.copy_(cache_pos)源码注释attention.py 第 311-313 行指出在 eager 模式下这个谓词会被特化specialize而在 export 后它成为SymBool即分支选择真正由张量数据驱动。这与参考实现形成清晰的对应关系参考实现传yNone导出版传一个全 NaN 张量两者结果一致——这一点被测试test_attention_torch_cond_eager显式覆盖第一次 forward 传真实的self.x第二次分别传empty_y全 NaN导出版与None参考版断言assert_close通过。当self.kv_cache is None时则退回普通路径y必须提供否则assert报错直接调用calculate_kv(y)。2.2 构造参数与校验MultiHeadAttention.__init__的签名与参考模块一致embed_dim、num_heads、num_kv_heads、head_dim、四个投影层q_proj/k_proj/v_proj/output_proj以及可选的pos_embeddings、q_norm/k_normQK-norm二者必须成对设置否则ValueError、kv_cache、max_seq_len默认 4096用于计算 RoPE 缓存、is_causal、attn_dropout。校验逻辑包括num_heads % num_kv_heads 0、embed_dim % num_heads 0、0 attn_dropout 1。一个实现细节SDPA 子模块构造时attn_dropout取self.attn_dropout if self.training else 0.0attention.py 第 164 行即评估/导出态下 dropout 恒为 0。2.3 SDPA 子模块head 维度对齐与 is_causal 逻辑SDPA.forwardattention.py 第 352-385 行接收[b, s, n_h, h_d]布局的 q/k/v先transpose(1, 2)到 SDPA 期望的[b, n, s, h_d]布局GQA 场景下通过unsqueeze(2).expand(...).flatten(1, 2)把 k/v 的num_kv_heads扩展到与 q 相同的num_heads调用底层 attention 函数由_sdpa_or_flex_attention()决定是 SDPA 还是 flex attention后者用于 sample packing 的 BlockMask 场景最后转回[b, s, -1]布局。is_causal的实际生效条件是一行合取is_causalself.kv_cache is None and mask is None and self.is_causal,即只有在无 KV cache、无显式 mask、且构造时is_causalTrue时才启用因果 mask——带 KV cache 的解码步需要显式传入因果 mask测试中即按self.causal_mask[input_pos, :]切片构造。2.4 setup_cache / reset_cachesetup_cache(batch_size, dtype, max_seq_len)若构造时已注入kv_cache则跳过并告警否则创建导出版InferenceKVCache注意此处transpose_cacheFalse并把cache_enabled置为Truereset_cache()未建缓存时抛RuntimeError否则调用kv_cache.reset()。cache_enabled标志的语义在源码注释中说明它指示 forward 期间是否更新缓存禁用时即使缓存已建立也走普通前向。三、导出版 KVCachetranspose_cache 与 clonekv_cache.py 中的KVCache继承自 torchtune/modules/kv_cache.py 的参考实现主要差异新增transpose_cache参数参考实现固定为转置布局(batch, n_kv, seq, h_d)transpose_cacheTrue时缓存形状为(batch_size, num_kv_heads, max_seq_len, head_dim)update()按k_out[:, :, cache_pos[:seq_len]] k_val写入transpose_cacheFalse时形状为(batch_size, max_seq_len, num_kv_heads, head_dim)update()按k_out[:, cache_pos[:seq_len]] k_val写入——与投影层输出的[b, s, n_kv, h_d]布局一致省去一次转置新增clone()方法创建同配置的新缓存并copy_三个 bufferk_cache/v_cache/cache_pos。这是torch.cond改写的配套——两个分支各自在克隆上操作互不污染输入位置追踪保持 compile-friendlyupdate()末尾self.cache_pos.add_(seq_len)kv_cache.py 第 130 行cache_pos从(0, 1, 2, 3, ...)每次前移seq_len用整型 buffer 而非 Pythonint记录当前写入位置避免引入动态性dynamism。assert (self.cache_pos[0] seq_len) self.max_seq_len保证不越界batch 超限时抛ValueError。forward中的copy_回写、clone()、setup_cache里transpose_cacheFalse的选择共同构成分支内只操作克隆 → 分支外合并回真实缓存的导出安全模式。四、导出版视觉位置编码_position_embeddings.py 提供 CLIP 视觉编码器所用位置编码的导出友好版本文件头注释说明其参考来源为 torchtune/models/clip/_position_embeddings.py核心改动是添加torch._check()以强制 symint 上的守卫guards在导出后仍被检查。4.1 TilePositionalEmbedding参数max_num_tiles图像最多可切分的 tile 数、embed_dimforward(x, aspect_ratio)x形状(bsz * n_imgs, n_tiles, n_tokens, embed_dim)aspect_ratio形状(bsz * n_imgs, 2)表示每张图 tile 裁剪前的行列 tile 数如(2, 1)。batch 内图像会 pad 到相同 tile 数位置编码只加到非 pad 的 tile 上导出友好细节第 187-219 行torch._check(n_tiles self.max_num_tiles)以及对n_tiles_h/n_tiles_w的torch._check_is_size、 1、 max_num_tiles等一系列守卫对切出的pos_embed先clone()再reshape——源码注释明确写道We need to do a clone here in order to make this model export friendly as the reshape is collapsing dim 0 and dim 1 into a single dimreshape 折叠维度前必须 clone否则导出失败加载时行为注册_load_state_dict_pre_hook在 ckpt 的max_num_tiles与实例化设置不一致时用F.interpolate(modebilinear, align_cornersTrue)插值 tile 维位置编码若参数是DTensor先full_tensor()收集完整张量、插值后再distribute_tensor重新分片。4.2 TiledTokenPositionalEmbedding用于每个 tile 不同、每个 token 也不同的 token 级位置编码包含两个参数local_token_positional_embedding(n_tokens_per_tile, embed_dim)每个 tile 相同、每个 token 不同global_token_positional_embedding(max_num_tiles, max_num_tiles, n_tokens_per_tile, embed_dim)tile 与 token 都不同两者按gate.tanh()做门控互补local 部分乘(1 - gate.tanh())global 部分乘gate.tanh()n_tokens_per_tile patch_grid_size**2 11 为 CLS token_load_state_dict_pre_hook分别对 local/global 做双线性插值且CLS token 不参与插值先切出[[0]]行插值图像 token 后再torch.cat拼回同样的导出友好处理torch._check守卫族 切块后clone()再reshape。4.3 批量替换函数与 MHA 的替换函数配套该文件提供两个递归替换工具用于把构建好的模型中参考版位置编码原地换成导出版并load_state_dict迁移参数replace_tile_positional_embedding(model)匹配torchtune.models.clip._position_embeddings.TilePositionalEmbeddingreplace_tiled_token_positional_embedding(model)匹配TiledTokenPositionalEmbedding从global_token_positional_embedding形状反推tile_sizeint(sqrt(n_tokens_per_tile - 1))、patch_size1。五、MHA 的一键替换replace_mha_with_inference_mhaattention.py 第 416-422 行 的公开入口replace_mha_with_inference_mha(module)递归遍历module.named_children()把每个参考版TorchTuneAttention.MultiHeadAttention替换为导出版逐项搬运embed_dim/num_heads/num_kv_heads/head_dim/四个投影层/pos_embeddings/q_norm/k_norm/kv_cache/max_seq_len/is_causal/attn_dropout投影层与位置编码直接复用同一批子模块对象权重零拷贝迁移from torchtune.modules._export.attention import replace_mha_with_inference_mha import torchtune.models.llama3_1 as llama3_1 model llama3_1.llama3_1_8b() model replace_mha_with_inference_mha(model) # 所有层 MHA 变为导出版六、等价性验证eager / export / AOTI 三层测试tests/torchtune/modules/_export/ 下的测试是该目录契约的执行保障其中 test_attention.py 构造了参考版与导出版共享同一批投影层、同一Llama3ScaledRoPE位置编码rope_base500_000scale_factor32的两个 MHAembed_dim20488 头并以参考版输出为基准做assert_closetest_attention_eager无缓存、setup_cache后首写缓存、带input_pos的读缓存input_pos[10..19]三类场景下 eager 数值一致test_attention_export用torch.export.export(..., strictTrue)导出且带动态序列长度seq_len_dim torch.export.Dim(seq_len, min1, max100) dynamic_shapes ( {0: torch.export.Dim.STATIC, 1: seq_len_dim, 2: torch.export.Dim.STATIC}, {0: torch.export.Dim.STATIC, 1: seq_len_dim, 2: torch.export.Dim.STATIC}, {0: torch.export.Dim.STATIC, 1: seq_len_dim}, ) ep torch.export.export(self.et_mha, (self.x, self.x), kwargs{input_pos: self.input_pos}, dynamic_shapesself.dynamic_shapes, strictTrue)导出的ep.module()输出与参考版 eager 输出一致test_attention_aotitorch._export.aot_inductor风格的完整 AOT 链路——torch._export.aot_compile(..., options{aot_inductor.package: True, reorder_for_peak_memory: False})经package_aoti打包成.pt2包后load_package加载执行结果再与参考版对齐test_attention_torch_cond_eager专测torch.cond改写等价性——参考版第二次解码传yNone导出版传全 NaN 的empty_y两者配合按input_pos切片的因果 mask输出assert_close。位置编码的测试在 test_export_position_embeddings.py 中README 要求这些测试在每日调度与触及_export目录的 PR 上运行。七、使用边界与注意事项结合 README 声明与源码、测试证据使用_export模块时应注意版本下限torch.cond路径要求 PyTorch ≥ 2.6.0export/AOTI 用例还依赖2.6.0.dev20241117之后的修复环境脚本钉在torch2.6.0接口承诺但 API 不稳定四个契约保证与参考模块同参、同输出、可导出但 README 明确 These modules are subject to change so proceed with caution跨版本升级后应重跑仓库内测试回归NaN 约定是 API 的一部分解码阶段读缓存必须传全 NaN 的y且需要显式传入按input_pos切好的因果 mask这不是 bug 而是torch.cond化的设计KV cache 布局setup_cache创建的缓存为transpose_cacheFalse[b, s, n_kv, h_d]若自行构造InferenceKVCache且设transpose_cacheTrueupdate()期望的 k/v 布局是[B, H, S, D]两者不可混用替换函数只改结构不训练replace_mha_with_inference_mha/replace_tile_positional_embedding系列仅用于把已构建可加载 ckpt的模型切换到导出路径权重通过复用子模块或load_state_dict完整保留。小结torchtune/modules/_export/的价值在于给出了训练/推理参考模块 → 可导出模块的一套可复用改写范式用torch.cond消除数据依赖的 Python 分支、用克隆-合并模式消除torch.cond分支内的输入变更、用独立 SDPA 子模块预留算子替换点、用 buffer 化cache_pos与torch._check守卫消除隐式动态性并以 eager/export/AOTI 三层数值等价测试锁定行为。这套模式对任何想把 TorchTune 构建的模型送入torch.export/AOT 编译管线如打包为.pt2部署的场景都可直接参照。【免费下载链接】torchtunePyTorch native post-training library项目地址: https://gitcode.com/GitHub_Trending/to/torchtune创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考