Python中如何用2位内存存储特定数值优化DQN经验回放缓冲区
我正在用Python编写Deep Q Network代码,电脑有32GB内存,但训练时经验回放缓冲区占满了RAM,遇到了大问题。我检查了经验回放缓冲区的代码,想找内存优化的空间——这个缓冲区存了两个各有100万个元素的numpy数组,数据类型是numpy.int8。但其中一个数组的元素只能是0、1、2、3,另一个只能是-1、0、1,这两个数组的每个元素其实只要2位内存就能存下。怎么创建每个元素占2位而不是8位的数组?我能接受一定程度的硬编码,比如类似下面这样的代码:
if bitarray[i][0] == 0 and bitarray[i][1] == 0: numberAtPositionI = -1
内存优化方案:用2位存储元素
方案1:用bitarray库实现(最简便)
bitarray库支持直接操作比特级数组,能轻松把每个元素压缩到2位,存取时手动做数值和比特的映射即可。
先安装库:
pip install bitarray
处理仅存0、1、2、3的数组
直接把数值转成2位二进制存入比特数组:
from bitarray import bitarray # 初始化比特数组:100万元素×2位 = 200万比特 arr_0_3 = bitarray(2 * 10**6) # 写入:把数值n(0-3)写入第i个位置 def write_0_3(arr, idx, n): start = idx * 2 # 把n转成2位二进制字符串,反转后存入(保证比特位顺序和读取逻辑匹配) bits = bin(n)[2:].zfill(2)[::-1] arr[start:start+2] = bitarray(bits) # 读取:从第i个位置取出数值 def read_0_3(arr, idx): start = idx * 2 bits = arr[start:start+2].to01()[::-1] return int(bits, 2) # 测试 write_0_3(arr_0_3, 0, 3) print(read_0_3(arr_0_3, 0)) # 输出3
处理仅存-1、0、1的数组
给每个数值分配固定2位编码:-1→00、0→01、1→10,剩下的11可留作异常处理:
arr_neg1_0_1 = bitarray(2 * 10**6) # 写入 def write_neg1_0_1(arr, idx, n): start = idx * 2 if n == -1: bits = '00' elif n == 0: bits = '01' elif n == 1: bits = '10' else: raise ValueError("仅支持-1、0、1") arr[start:start+2] = bitarray(bits) # 读取 def read_neg1_0_1(arr, idx): start = idx * 2 bits = arr[start:start+2].to01() if bits == '00': return -1 elif bits == '01': return 0 elif bits == '10': return 1 else: return 0 # 处理异常情况 # 测试 write_neg1_0_1(arr_neg1_0_1, 0, -1) write_neg1_0_1(arr_neg1_0_1, 1, 1) print(read_neg1_0_1(arr_neg1_0_1, 0)) # 输出-1 print(read_neg1_0_1(arr_neg1_0_1, 1)) # 输出1
方案2:手动用numpy uint8数组打包(无额外依赖)
每个uint8字节能存4个2位元素(8÷2=4),100万元素只需要25万个uint8,内存直接降到原来的1/4,完全不需要额外库。
处理0-3的数组
import numpy as np # 初始化打包数组:向上取整100万÷4,得到250000个uint8元素 pack_size = (10**6 + 3) // 4 packed_0_3 = np.zeros(pack_size, dtype=np.uint8) # 打包写入 def pack_0_3(idx, n): byte_idx = idx // 4 bit_pos = idx % 4 shift = bit_pos * 2 # 先清空目标位置的比特,再写入数值 packed_0_3[byte_idx] &= ~(0b11 << shift) packed_0_3[byte_idx] |= (n & 0b11) << shift # 解包读取 def unpack_0_3(idx): byte_idx = idx // 4 bit_pos = idx % 4 shift = bit_pos * 2 return (packed_0_3[byte_idx] >> shift) & 0b11 # 测试 pack_0_3(0, 3) pack_0_3(1, 1) print(unpack_0_3(0)) # 输出3 print(unpack_0_3(1)) # 输出1
处理-1、0、1的数组
先把数值映射成0-2的编码(-1→0、0→1、1→2),再用上面的方法打包:
packed_neg1_0_1 = np.zeros(pack_size, dtype=np.uint8) # 打包写入 def pack_neg1_0_1(idx, n): # 映射成0-2的编码 if n == -1: code = 0 elif n == 0: code = 1 elif n == 1: code = 2 else: raise ValueError("仅支持-1、0、1") byte_idx = idx // 4 bit_pos = idx % 4 shift = bit_pos * 2 packed_neg1_0_1[byte_idx] &= ~(0b11 << shift) packed_neg1_0_1[byte_idx] |= (code & 0b11) << shift # 解包读取 def unpack_neg1_0_1(idx): byte_idx = idx // 4 bit_pos = idx % 4 shift = bit_pos * 2 code = (packed_neg1_0_1[byte_idx] >> shift) & 0b11 # 映射回原数值 if code == 0: return -1 elif code == 1: return 0 elif code == 2: return 1 else: return 0 # 测试 pack_neg1_0_1(0, -1) pack_neg1_0_1(2, 1) print(unpack_neg1_0_1(0)) # 输出-1 print(unpack_neg1_0_1(2)) # 输出1
额外说明
两种方法都会增加少量打包/解包的计算开销,但对于DQN训练来说,这点开销远小于内存不足导致的卡顿或崩溃。如果需要频繁随机存取,bitarray的速度会略快一些。
内容的提问来源于stack exchange,提问作者Bryan Carty
相关产品推荐
相关产品推荐

