TF-Agents自定义PyEnvironment中如何将观测值设为字符串类型
TF-Agents官方环境教程提供了一个受二十一点玩法启发的卡牌游戏环境示例,其初始化代码中使用BoundedArraySpec定义数值型的动作、观测规范:
class CardGameEnv(py_environment.PyEnvironment): def __init__(self): self._action_spec = array_spec.BoundedArraySpec( shape=(), dtype=np.int32, minimum=0, maximum=1, name='action') self._observation_spec = array_spec.BoundedArraySpec( shape=(1,), dtype=np.int32, minimum=0, name='observation') self._state = 0 self._episode_ended = False
当业务场景需要使用字符串类型观测值时,直接给BoundedArraySpec传入dtype=str会触发类型校验错误:
TypeError: Cannot find minimum value of <dtype: 'string'> with type <dtype: 'string'>. In call to configurable 'BoundedArraySpec' (<class 'tf_agents.specs.array_spec.BoundedArraySpec'>)
报错原因很直接:BoundedArraySpec的设计目标就是定义有明确取值上下界的数值型数组,字符串类型没有通用的可比较大小的上下界,自然过不了校验。
可落地的解决方法
1. 替换为无界的ArraySpec
不需要做编码转换的最轻量方案,直接用不需要传入minimum/maximum参数的基础ArraySpec定义观测规范即可,注意要使用numpy的字符串类型而非Python原生str:
self._observation_spec = array_spec.ArraySpec( shape=(1,), dtype=np.str_, name='observation' )
这个方案的缺点是部分旧版本TF-Agents的ReplayBuffer、TF环境包装器可能对字符串dtype兼容不好,适合快速做原型验证用。
2. 将字符串编码为数值(生产环境首选)
强化学习训练链路本身就是基于数值张量运行的,把字符串观测做映射编码是兼容性最好、运行效率最高的方案:
- 提前枚举所有可能出现的字符串观测值,构建字符串到整数的一一映射表,比如卡牌场景下
{"红桃A": 0, "红桃2":1, ..., "黑桃K":51} - 环境内部状态存储编码后的整数值,观测规范照常使用
BoundedArraySpec,minimum设为0,maximum设为映射表的长度减1即可 - 如果调试时需要查看原始字符串观测,只需要在
_reset、_step方法返回TimeStep前加一步临时解码逻辑,训练过程全程跑数值,不会遇到任何类型兼容问题。
如果观测字符串是开放集合(无法提前枚举全部值),可以用固定长度哈希、字符嵌入等方式把字符串转成固定维度的数值数组,再定义对应shape的数值型观测规范即可。
3. 自定义Spec类
如果必须在观测中原样传递字符串,且当前版本ArraySpec不支持字符串类型校验,可以继承ArraySpec基类自定义一个字符串专用的Spec,重写其中的取值校验逻辑,跳过min/max边界检查即可。这个方案灵活度最高,但需要自行适配后续TF张量转换、经验回放存储等环节的逻辑,适合有定制化开发需求的场景。
内容的提问来源于stack exchange,提问作者flowerboy

