DDPG经验回放:优雅提取状态与下一状态的NumPy实现问询
嘿,我在实现DDPG的时候也踩过这个坑!当时用deque存经验元组,采样后手动解包总觉得代码又丑又冗余,后来摸索出几个优雅的办法,帮你解决这个问题:
为什么np.hsplit没用?
先给你捋清楚原因:np.hsplit是用来对单个numpy数组按列进行分割的,但你采样得到的是元组列表,每个元组里的state、next_state都是独立的数组,并不是一个大数组的列,所以hsplit自然不适用啦。
方法1:用zip(*samples)批量解包(强推!)
这是Python处理这类元组列表最简洁高效的方式,没有之一。zip(*samples)会把所有元组对应位置的元素打包到一起,完美匹配你需要的state、action、reward、next_state分组。
看代码示例:
import numpy as np from collections import deque # 模拟经验池填充 experience_pool = deque() # 假设state/next_state是4维向量,action是1维,reward是标量 for _ in range(20): state = np.random.rand(4) action = np.random.rand(1) reward = np.random.uniform(-1, 1) next_state = np.random.rand(4) experience_pool.append((state, action, reward, next_state)) # 采样一批经验 batch_size = 8 samples = [experience_pool[i] for i in np.random.choice(len(experience_pool), batch_size)] # 核心解包操作! states, actions, rewards, next_states = zip(*samples) # 转成numpy数组(DDPG里后续网络输入需要数组形式) states = np.array(states) next_states = np.array(next_states) print(states.shape) # 输出 (8, 4),符合批量输入的格式 print(next_states.shape) # 同样是 (8, 4)
原理很简单:*samples会把列表里的每个元组拆成独立参数传给zip,zip就会依次抓取每个元组的第0个、第1个...元素,打包成新的元组,直接赋值给四个变量就行,完全不用手动遍历索引。
方法2:列表推导式(适合需要额外处理的场景)
如果你需要对state或者next_state做一些即时处理(比如归一化、裁剪),列表推导式会更直观:
# 提取state并转数组 states = np.array([sample[0] for sample in samples]) # 提取next_state并转数组 next_states = np.array([sample[3] for sample in samples])
这种写法虽然不如zip简洁胜在灵活,比如可以直接在推导式里加处理逻辑:
# 提取state并做归一化 states = np.array([sample[0] / 10.0 for sample in samples])
总结
优先用zip(*samples)的方式,代码简洁、执行高效,完全符合Python的惯用写法,能让你的DDPG经验回放部分看起来清爽很多~
内容的提问来源于stack exchange,提问作者hakaishinbeerus
相关产品推荐
相关产品推荐

