You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 02:05:45