使用普通数组索引Numpy ndarray时为何抛出IndexError?
问题解决方法
核心问题1:多维数组索引方式错误
你用长度为4的数组直接索引5维q_table时,numpy会把这个数组当成对第一个轴的批量索引(即取q_table[arr[0], :, :, :, :]、q_table[arr[1], :, :, :, :]等),而不是分别对应前4个轴的单索引。这就导致如果数组里的元素超过第一个轴的大小(10),就会触发索引越界错误。
解决方法:把索引数组转换成tuple,这样numpy会将tuple中的每个元素对应到多维数组的各个轴上:
max_val = np.max(self.q_table[tuple(self.quantize_state(observation_space, [-150, 100, 3, 3]))])
核心问题2:静态方法无法访问局部变量
你的quantize_state是静态方法,无法访问__init__里定义的局部变量OBSERVATION_SPACE_RESOLUTION,这会导致生成的分割点数量错误,最终让digitize返回超出维度范围的索引(比如错误得到11,而第一个轴只有10个元素)。
解决方法:把OBSERVATION_SPACE_RESOLUTION设为类属性,让静态方法可以正确访问:
class AgentBase: # 把分辨率设为类属性,让静态方法能访问 OBSERVATION_SPACE_RESOLUTION = [10, 15, 15, 15] def __init__(self, observation_space): self.q_table = np.zeros([*self.OBSERVATION_SPACE_RESOLUTION, 4]) max_val = np.max(self.q_table[tuple(self.quantize_state(observation_space, [-150, 100, 3, 3]))]) print(max_val) @staticmethod def quantize_state(observation_space, state): state_quantized = np.zeros(len(state)) lin_spaces = [] for i in range(len(observation_space)): # 访问类属性获取分辨率 lin_spaces.append(np.linspace(observation_space[i][0], observation_space[i][1], AgentBase.OBSERVATION_SPACE_RESOLUTION[i] - 1, dtype=int)) for i in range(len(lin_spaces)): state_quantized[i] = np.digitize(state[i], lin_spaces[i]) return state_quantized.astype(int)
额外验证点
检查digitize的返回值范围:digitize会返回0到len(lin_space)的整数,而对应轴的大小是OBSERVATION_SPACE_RESOLUTION[i],len(lin_space) = OBSERVATION_SPACE_RESOLUTION[i]-1,所以返回的最大索引刚好是轴的合法最大索引(0到size-1),不会越界。只要保证OBSERVATION_SPACE_RESOLUTION的取值一致,这部分就没问题。
内容的提问来源于stack exchange,提问作者Olivér Raisz
相关产品推荐
相关产品推荐

