首页
/ OpenSpiel项目中RCFR算法在Keras 3下的兼容性问题分析

OpenSpiel项目中RCFR算法在Keras 3下的兼容性问题分析

2025-06-13 07:47:54作者:柏廷章Berta

问题背景

OpenSpiel是一个由Google DeepMind开发的开源游戏AI研究平台,其中包含了多种强化学习算法实现。RCFR(Regression Counterfactual Regret Minimization)是其中一种重要的算法实现,它结合了传统的CFR算法与神经网络回归技术。

近期在Ubuntu 24.04系统下,使用Python 3.12和Keras 3.1.1运行RCFR测试时出现了兼容性问题。本文将详细分析这些问题及其解决方案。

主要错误现象

测试过程中出现了多个关键错误,主要集中在以下几个方面:

  1. 优化器参数错误:测试中尝试使用tf.keras.optimizers.Adam(lr=0.005, amsgrad=True)时,Keras 3抛出了参数不识别错误,提示lr参数不被接受。

  2. 数组转换警告:多处出现了NumPy数组转换为标量的DeprecationWarning,提示在NumPy 1.25及以后版本中,这种转换将会报错。

  3. 函数重追踪警告:出现了关于tf.function频繁重追踪的性能警告。

问题根源分析

Keras 3 API变更

Keras 3对优化器API进行了重大调整,最显著的变化是:

  • lr参数已更名为learning_rate,这是导致测试失败的直接原因
  • 参数验证更加严格,不再接受旧版参数名称
  • 内部实现机制有所变化,可能导致其他潜在兼容性问题

NumPy版本兼容性

测试中出现的数组转换警告反映了NumPy 1.25版本对数组处理方式的变更:

  • 不再允许直接将多维数组隐式转换为标量
  • 需要显式提取单个元素后再进行标量操作

TensorFlow函数优化

频繁的函数重追踪警告表明:

  • 在循环中重复创建@tf.function装饰的函数
  • 可能传递了形状不一致的张量
  • 或者传递了Python对象而非张量

解决方案建议

优化器参数修正

将所有的lr=参数替换为learning_rate=,例如:

# 旧代码
optimizer = tf.keras.optimizers.Adam(lr=0.005, amsgrad=True)

# 新代码
optimizer = tf.keras.optimizers.Adam(learning_rate=0.005, amsgrad=True)

数组处理规范化

对于NumPy数组转换问题,需要显式提取元素:

# 旧代码
reach_probabilities[player] = next_reach_prob

# 新代码
reach_probabilities[player] = next_reach_prob.item()  # 显式转换为Python标量

函数优化建议

对于函数重追踪问题,可以:

  1. @tf.function装饰器移到循环外部
  2. 确保传递的张量形状一致
  3. 使用reduce_retracing=True选项减少不必要的重追踪

实施验证

在实际修复过程中,需要注意:

  1. 全面检查所有优化器实例化代码
  2. 对数组操作进行彻底审查
  3. 测试不同游戏场景下的算法表现
  4. 监控训练过程中的性能指标

结论

Keras 3的API变更带来了必要的现代化改进,但也需要相应的代码适配。通过系统性地解决参数命名、数组处理和函数优化等问题,可以确保RCFR算法在新版本框架下的稳定运行。这类兼容性问题在深度学习框架升级过程中较为常见,理解其背后的设计变更有助于更好地维护和升级算法实现。

登录后查看全文
热门项目推荐

项目优选

收起
docsdocs
OpenHarmony documentation | OpenHarmony开发者文档
Dockerfile
152
1.97 K
kernelkernel
deepin linux kernel
C
22
6
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
494
37
communitycommunity
本项目是CANN开源社区的核心管理仓库,包含社区的治理章程、治理组织、通用操作指引及流程规范等基础信息
323
10
openGauss-serveropenGauss-server
openGauss kernel ~ openGauss is an open source relational database management system
C++
145
191
openHiTLSopenHiTLS
旨在打造算法先进、性能卓越、高效敏捷、安全可靠的密码套件,通过轻量级、可剪裁的软件技术架构满足各行业不同场景的多样化要求,让密码技术应用更简单,同时探索后量子等先进算法创新实践,构建密码前沿技术底座!
C
991
395
nop-entropynop-entropy
Nop Platform 2.0是基于可逆计算理论实现的采用面向语言编程范式的新一代低代码开发平台,包含基于全新原理从零开始研发的GraphQL引擎、ORM引擎、工作流引擎、报表引擎、规则引擎、批处理引引擎等完整设计。nop-entropy是它的后端部分,采用java语言实现,可选择集成Spring框架或者Quarkus框架。中小企业可以免费商用
Java
8
0
ohos_react_nativeohos_react_native
React Native鸿蒙化仓库
C++
193
277
RuoYi-Vue3RuoYi-Vue3
🎉 (RuoYi)官方仓库 基于SpringBoot,Spring Security,JWT,Vue3 & Vite、Element Plus 的前后端分离权限管理系统
Vue
937
554
金融AI编程实战金融AI编程实战
为非计算机科班出身 (例如财经类高校金融学院) 同学量身定制,新手友好,让学生以亲身实践开源开发的方式,学会使用计算机自动化自己的科研/创新工作。案例以量化投资为主线,涉及 Bash、Python、SQL、BI、AI 等全技术栈,培养面向未来的数智化人才 (如数据工程师、数据分析师、数据科学家、数据决策者、量化投资人)。
Python
75
70