0

0

JAX vjp 失败原因与 vmap + custom_vjp 的正确组合方式

霞舞

霞舞

发布时间:2026-01-09 12:53:25

|

160人浏览过

|

来源于php中文网

原创

JAX vjp 失败原因与 vmap + custom_vjp 的正确组合方式

当对带有 `custom_vjp` 的函数先 `vmap` 再调用 `vjp` 时,若在定义 `vmap` 版本后覆盖了原始函数名,会导致前向传播中递归调用错误的 vmapped 版本,从而引发 cotangent 形状不匹配的错误。

在 JAX 中,custom_vjp 的前向函数(fwd)必须严格调用原始未变换的函数,以确保其输入/输出形状与 vjp 约定一致:即前向传播返回的 primal_out 形状应与后续 vjp 接收的 cotangent 形状完全匹配(即 cotangent.shape == primal_out.shape)。

问题代码中,关键错误在于:

test_func = vmap(test_func, in_axes=(None, 0))  # ❌ 覆盖了原始 test_func

这导致 test_func_fwd 内部调用的 test_func(jnp.dot(R, R)) 实际执行的是 已 vmapped 的版本,而 jnp.dot(R, R) 的输入 R 是标量(因 R 是 jnp.dot 的结果,shape 为 ()),但 vmapped test_func 期望 R 具有 batch 维度(如 (10, 3)),于是内部逻辑错乱,最终使 primal_out 的隐式形状与 vjp 期望不符——vjp 认为输出是 (10,),但 bwd 接收到的 residual 和 cotangent 却因前向误调而维度失配,触发报错:

ValueError: Shape of cotangent input to vjp pullback function (10,) must be the same as the shape of corresponding primal input (10, 3).

该错误信息虽表述为“cotangent 应与 primal input 同形”,实则是 JAX 在反向传播校验阶段,因前向路径被污染,无法正确推导出梯度传播所需的张量结构所致。

火山方舟
火山方舟

火山引擎一站式大模型服务平台,已接入满血版DeepSeek

下载

✅ 正确做法是:保留原始 test_func 不变,仅将 vmap 结果赋给新变量名

# ✅ 保持原始 test_func 不被覆盖
test_func_mapped = vmap(test_func, in_axes=(None, 0))

# 在 vjp 中使用映射后的版本
primal, f_vjp = vjp(partial(test_func_mapped, f), jnp.ones((10, 3)))
cotangent = jnp.ones(10)  # shape matches primal_out: (10,)
cotangent_out = f_vjp(cotangent)

print(cotangent_out[0].shape)  # → (10, 3)

? 补充注意事项:

  • custom_vjp 的 fwd 函数中禁止调用任何高阶变换(如 vmap, jit, grad)后的函数,除非明确设计为支持嵌套;
  • 若需批量处理并保留 vjp 可用性,推荐使用 vmap 包裹整个 vjp 调用(即 vmap(vjp(...))),而非先 vmap 函数再 vjp;
  • 对于复杂控制流或状态依赖场景,建议通过 jax.custom_vjp + non-differentiable arguments 显式隔离不变参数,并始终在 fwd/bwd 中使用原始函数引用。

遵循“函数变换不覆盖原名”这一原则,可避免绝大多数 vmap 与 custom_vjp 组合时的静默行为异常。

热门AI工具

更多
DeepSeek
DeepSeek

幻方量化公司旗下的开源大模型平台

豆包大模型
豆包大模型

字节跳动自主研发的一系列大型语言模型

通义千问
通义千问

阿里巴巴推出的全能AI助手

腾讯元宝
腾讯元宝

腾讯混元平台推出的AI助手

文心一言
文心一言

文心一言是百度开发的AI聊天机器人,通过对话可以生成各种形式的内容。

讯飞写作
讯飞写作

基于讯飞星火大模型的AI写作工具,可以快速生成新闻稿件、品宣文案、工作总结、心得体会等各种文文稿

即梦AI
即梦AI

一站式AI创作平台,免费AI图片和视频生成。

ChatGPT
ChatGPT

最最强大的AI聊天机器人程序,ChatGPT不单是聊天机器人,还能进行撰写邮件、视频脚本、文案、翻译、代码等任务。

相关专题

更多
点击input框没有光标怎么办
点击input框没有光标怎么办

点击input框没有光标的解决办法:1、确认输入框焦点;2、清除浏览器缓存;3、更新浏览器;4、使用JavaScript;5、检查硬件设备;6、检查输入框属性;7、调试JavaScript代码;8、检查页面其他元素;9、考虑浏览器兼容性。本专题为大家提供相关的文章、下载、课程内容,供大家免费下载体验。

196

2023.11.24

Swift iOS架构设计与MVVM模式实战
Swift iOS架构设计与MVVM模式实战

本专题聚焦 Swift 在 iOS 应用架构设计中的实践,系统讲解 MVVM 模式的核心思想、数据绑定机制、模块拆分策略以及组件化开发方法。内容涵盖网络层封装、状态管理、依赖注入与性能优化技巧。通过完整项目案例,帮助开发者构建结构清晰、可维护性强的 iOS 应用架构体系。

24

2026.03.03

C++高性能网络编程与Reactor模型实践
C++高性能网络编程与Reactor模型实践

本专题围绕 C++ 在高性能网络服务开发中的应用展开,深入讲解 Socket 编程、多路复用机制、Reactor 模型设计原理以及线程池协作策略。内容涵盖 epoll 实现机制、内存管理优化、连接管理策略与高并发场景下的性能调优方法。通过构建高并发网络服务器实战案例,帮助开发者掌握 C++ 在底层系统与网络通信领域的核心技术。

25

2026.03.03

Golang 测试体系与代码质量保障:工程级可靠性建设
Golang 测试体系与代码质量保障:工程级可靠性建设

Go语言测试体系与代码质量保障聚焦于构建工程级可靠性系统。本专题深入解析Go的测试工具链(如go test)、单元测试、集成测试及端到端测试实践,结合代码覆盖率分析、静态代码扫描(如go vet)和动态分析工具,建立全链路质量监控机制。通过自动化测试框架、持续集成(CI)流水线配置及代码审查规范,实现测试用例管理、缺陷追踪与质量门禁控制,确保代码健壮性与可维护性,为高可靠性工程系统提供质量保障。

77

2026.02.28

Golang 工程化架构设计:可维护与可演进系统构建
Golang 工程化架构设计:可维护与可演进系统构建

Go语言工程化架构设计专注于构建高可维护性、可演进的企业级系统。本专题深入探讨Go项目的目录结构设计、模块划分、依赖管理等核心架构原则,涵盖微服务架构、领域驱动设计(DDD)在Go中的实践应用。通过实战案例解析接口抽象、错误处理、配置管理、日志监控等关键工程化技术,帮助开发者掌握构建稳定、可扩展Go应用的最佳实践方法。

60

2026.02.28

Golang 性能分析与运行时机制:构建高性能程序
Golang 性能分析与运行时机制:构建高性能程序

Go语言以其高效的并发模型和优异的性能表现广泛应用于高并发、高性能场景。其运行时机制包括 Goroutine 调度、内存管理、垃圾回收等方面,深入理解这些机制有助于编写更高效稳定的程序。本专题将系统讲解 Golang 的性能分析工具使用、常见性能瓶颈定位及优化策略,并结合实际案例剖析 Go 程序的运行时行为,帮助开发者掌握构建高性能应用的关键技能。

48

2026.02.28

Golang 并发编程模型与工程实践:从语言特性到系统性能
Golang 并发编程模型与工程实践:从语言特性到系统性能

本专题系统讲解 Golang 并发编程模型,从语言级特性出发,深入理解 goroutine、channel 与调度机制。结合工程实践,分析并发设计模式、性能瓶颈与资源控制策略,帮助将并发能力有效转化为稳定、可扩展的系统性能优势。

26

2026.02.27

Golang 高级特性与最佳实践:提升代码艺术
Golang 高级特性与最佳实践:提升代码艺术

本专题深入剖析 Golang 的高级特性与工程级最佳实践,涵盖并发模型、内存管理、接口设计与错误处理策略。通过真实场景与代码对比,引导从“可运行”走向“高质量”,帮助构建高性能、可扩展、易维护的优雅 Go 代码体系。

20

2026.02.27

Golang 测试与调试专题:确保代码可靠性
Golang 测试与调试专题:确保代码可靠性

本专题聚焦 Golang 的测试与调试体系,系统讲解单元测试、表驱动测试、基准测试与覆盖率分析方法,并深入剖析调试工具与常见问题定位思路。通过实践示例,引导建立可验证、可回归的工程习惯,从而持续提升代码可靠性与可维护性。

4

2026.02.27

热门下载

更多
网站特效
/
网站源码
/
网站素材
/
前端模板

精品课程

更多
相关推荐
/
热门推荐
/
最新课程
Java 教程
Java 教程

共578课时 | 76.6万人学习

国外Web开发全栈课程全集
国外Web开发全栈课程全集

共12课时 | 1万人学习

关于我们 免责申明 举报中心 意见反馈 讲师合作 广告合作 最新更新
php中文网:公益在线php培训,帮助PHP学习者快速成长!
关注服务号 技术交流群
PHP中文网订阅号
每天精选资源文章推送

Copyright 2014-2026 https://www.php.cn/ All Rights Reserved | php.cn | 湘ICP备2023035733号