强化学习视频分类模型内存分配失败问题求助
问题概述
使用Stable Baselines3的PPO框架构建视频分类模型时,初始化阶段触发MemoryError,报错显示需要分配291GiB内存存储rollout buffer的观测数据(数组形状(2048,4,3,60,230,230)),超出硬件内存限制。已尝试调整数据加载的batch_size=5,但未解决核心内存占用问题。
可行优化方案
1. 缩小Rollout Buffer尺寸(核心优化)
Stable Baselines3中PPO的默认n_steps=2048是rollout buffer的容量,直接决定内存占用。将其大幅调小,比如设置为64或32:
# 初始化PPO时指定n_steps参数 model = PPO('CnnPolicy', env, policy_kwargs=policy_kwargs, verbose=1, n_steps=64)
内存计算:64步×4环境×3通道×60帧×230×230像素×4字节(float32)≈8.9GiB,普通设备可承受;若仍有压力,可降至32,内存占用减半。
2. 减少并行环境数量
当前使用n_envs=4并行环境,减少到1或2可进一步降低内存压力:
# 调整并行环境数量为1 env = make_vec_env(lambda: env, n_envs=1)
配合n_steps=64时,内存占用降至≈2.2GiB,完全满足普通硬件需求。
3. 压缩观测数据维度
原始观测(3,60,230,230)的维度过高,通过预处理降采样压缩:
- 添加预处理函数,对视频帧做空间和时间降采样,并归一化到[0,1]范围:
import cv2 def preprocess_frame(frame): # 转换维度方便cv2处理:(3,60,230,230) → (60,230,230,3) frame = np.transpose(frame, (1,2,3,0)) # 空间降采样:230×230 → 115×115 frame = np.array([cv2.resize(f, (115,115)) for f in frame]) # 时间降采样:60帧 → 30帧(每隔2帧取1帧) frame = frame[::2] # 转回原维度顺序:(60,115,115,3) → (3,30,115,115) frame = np.transpose(frame, (3,0,1,2)) return frame.astype(np.float32) / 255.0
- 修改环境的
reset和step方法,返回预处理后的帧:
def reset(self): self.current_index = 0 try: self.current_batch, self.current_labels = next(self.batch_gen) except StopIteration: self.batch_gen = batch_loader(self.ttset_training, self.batch_size) self.current_batch, self.current_labels = next(self.batch_gen) # 应用预处理 return preprocess_frame(self.current_batch[self.current_video][self.current_index]) def step(self, action): reward = 1 if action == self.current_labels[self.current_video][self.current_index] else -1 self.current_index += 1 done = False if self.current_index >= len(self.current_batch[self.current_video]): self.current_video += 1 self.current_index = 0 if self.current_video >= len(self.current_batch): done = True next_state = preprocess_frame(self.current_batch[self.current_video][self.current_index]) if not done else np.zeros((3,30,115,115), dtype=np.float32) return next_state, reward, done, {}
- 更新环境的观测空间定义:
self.observation_space = spaces.Box(low=0, high=1, shape=(3, 30, 115, 115), dtype=np.float32)
此优化可将观测数据的内存占用降至原来的1/8,配合前两项优化,内存压力会大幅降低。
4. 简化CNN模型的全连接层
原模型的全连接层fc1输入维度为64*30*115*115(约2500万参数),显存占用极高。改用全局平均池化层替代直接flatten:
class CNNModel(nn.Module): def __init__(self, num_classes): super(CNNModel, self).__init__() self.conv1 = nn.Conv3d(3, 32, kernel_size=3, stride=1, padding=1) self.conv2 = nn.Conv3d(32, 64, kernel_size=3, stride=1, padding=1) # 添加全局平均池化层 self.global_avg_pool = nn.AdaptiveAvgPool3d(1) self.fc1 = nn.Linear(64, 512) self.fc2 = nn.Linear(512, num_classes) def forward(self, x): x = F.relu(self.conv1(x)) x = F.max_pool3d(x, 2) x = F.relu(self.conv2(x)) x = F.max_pool3d(x, 2) x = self.global_avg_pool(x) x = x.view(-1, 64) # 池化后仅保留通道维度 x = F.relu(self.fc1(x)) x = self.fc2(x) return x
修改后fc1的输入维度从2500万降至64,模型参数数量大幅减少,显存占用显著降低。
5. 实现按需加载视频帧
避免一次性加载所有视频帧到内存,改为按需从磁盘读取:
- 修改
batch_loader返回数据索引而非加载好的帧:
def batch_loader(data_indices, batch_size): data_len = len(data_indices) indices = np.arange(data_len) np.random.shuffle(indices) for start_idx in range(0, data_len, batch_size): excerpt = indices[start_idx:start_idx + batch_size] yield [data_indices[i] for i in excerpt]
- 环境中添加帧加载逻辑,仅在需要时读取当前视频:
class VideoClassificationEnv(gym.Env): def __init__(self, data_indices, batch_size): super(VideoClassificationEnv, self).__init__() self.data_indices = data_indices # 存储视频路径/索引 self.batch_size = batch_size self.current_batch_indices = None self.current_video_frames = None self.current_video_labels = None self.current_index = 0 self.current_video = 0 self.observation_space = spaces.Box(low=0, high=1, shape=(3, 30, 115, 115), dtype=np.float32) self.action_space = spaces.Discrete(20) self.batch_gen = batch_loader(self.data_indices, self.batch_size) def load_video(self, idx): # 根据索引/路径加载视频帧和标签,替换为实际逻辑 frames = ... # 返回(3,60,230,230)数组 labels = ... # 返回对应标签数组 return frames, labels def reset(self): self.current_index = 0 try: self.current_batch_indices = next(self.batch_gen) except StopIteration: self.batch_gen = batch_loader(self.data_indices, self.batch_size) self.current_batch_indices = next(self.batch_gen) # 加载当前视频数据 self.current_video_frames, self.current_video_labels = self.load_video(self.current_batch_indices[self.current_video]) return preprocess_frame(self.current_video_frames[self.current_index]) # step方法对应修改为使用self.current_video_frames和self.current_video_labels
此优化可避免原始视频帧长期占用内存,仅当前批次的视频数据会被加载,进一步释放内存资源。
内容的提问来源于stack exchange,提问作者Paarth Jha

