基于PPO的视频分类RL模型训练CUDA显存不足问题求助
解决PPO训练视频分类RL模型的CUDA显存不足问题
问题背景
我用Stable Baselines3的PPO算法训练视频分类强化学习模型,输入视频张量frame_tensors维度为1927x1x3x60x230x230,标签labels维度为1927x1。训练时出现CUDA显存不足报错,使用的GPU是RTX 4070Ti(12GB显存),已将batch size降至4且无法缩小帧尺寸(会丢失关键信息),需要可行的解决方案。
训练代码
from TTset_main import load_TTset # File paths for training and validation data filted_tset_path = "path.csv" filted_vset_path = "path.csv" # Load training and validation sets TTset_training, TTset_vali, training_data_length = load_TTset(filted_tset_path, filted_vset_path) # Extract frame tensors and labels from the training set frame_tensors = [] labels = [] for inputs, targets in TTset_training: frame_tensors.append(inputs) labels.append(targets) # Print dimensions of the frame tensors and labels print("Inputs") print("Dimension 1: ", len(frame_tensors)) #1927 print("Dimension 2: ", len(frame_tensors[0])) #1 print("Dimension 3: ", len(frame_tensors[0][0])) #3 print("Dimension 4: ", len(frame_tensors[0][0][0])) #60 print("Dimension 5: ", len(frame_tensors[0][0][0][0])) #230 print("Dimension 6: ", len(frame_tensors[0][0][0][0][0])) #230 print() print("Training data length:", training_data_length) print("Targets") print("Dimension 1: ", len(labels)) #1927 print("Dimension 2: ", len(labels[0])) #1 import numpy as np import gymnasium as gym from gymnasium import spaces from stable_baselines3 import PPO from stable_baselines3.common.vec_env import DummyVecEnv from stable_baselines3.common.callbacks import BaseCallback from tqdm import tqdm # Custom gym environment for table tennis class TableTennisEnv(gym.Env): def __init__(self, frame_tensors, labels, frame_size=(3, 60, 230, 230)): super(TableTennisEnv, self).__init__() self.frame_tensors = frame_tensors self.labels = labels self.current_step = 0 self.current_substep = 0 self.frame_size = frame_size self.n_actions = 20 # Number of unique actions self.observation_space = spaces.Box(low=0, high=255, shape=(3, 60, 230, 230), dtype=np.float32) self.action_space = spaces.Discrete(self.n_actions) self.normalize_images = False def reset(self, seed=None): self.current_step = 0 self.current_substep = 0 return self.frame_tensors[self.current_step][self.current_substep], {} def step(self, action): reward = 1 if action == self.labels[self.current_step][self.current_substep] else -1 self.current_substep += 1 if self.current_substep >= len(self.frame_tensors[self.current_step]): self.current_substep = 0 self.current_step += 1 done = self.current_step >= len(frame_tensors) obs = self.frame_tensors[self.current_step][self.current_substep] if not done else np.zeros_like(frame_tensors[0][0]) return obs, reward, done, {} def render(self, mode='human'): pass # Reduce memory usage by processing in smaller batches env = DummyVecEnv([lambda: TableTennisEnv(frame_tensors, labels, frame_size=(3, 60, 230, 230))]) # Callback for progress bar during training class ProgressBarCallback(BaseCallback): def __init__(self, total_timesteps, verbose=0): super(ProgressBarCallback, self).__init__(verbose) self.total_timesteps = total_timesteps self.pbar = None def _on_training_start(self): self.pbar = tqdm(total=self.total_timesteps) def _on_step(self): self.pbar.update(self.model.n_steps) return True def _on_training_end(self): self.pbar.close() # Set total timesteps for training total_timesteps = 2048 # Adjust timesteps as needed # Initialize PPO model with a smaller batch size model = PPO("MlpPolicy", env, verbose=1, n_steps=64, tensorboard_log="./ppo_tt_tensorboard/") # Reduce batch size progress_bar_callback = ProgressBarCallback(total_timesteps=total_timesteps) # Train the model model.learn(total_timesteps=total_timesteps, callback=progress_bar_callback) # Save the trained model model.save("ppo_table_tennis")
报错信息
Traceback (most recent call last): File "d:/PJ/TTARnet_v2_to_lab_兆仁code/prototype.py", line 107, in <module> model.learn(total_timesteps=total_timesteps, callback=progress_bar_callback) File "C:\Users\User1\anaconda3\envs\ML\lib\site-packages\stable_baselines3\ppo\ppo.py", line 315, in learn return super().learn( File "C:\Users\User1\anaconda3\envs\ML\lib\site-packages\stable_baselines3\common\on_policy_algorithm.py", line 313, in learn self.train() File "C:\Users\User1\anaconda3\envs\ML\lib\site-packages\stable_baselines3\ppo\ppo.py", line 282, in train self.policy.optimizer.step() File "C:\Users\User1\anaconda3\envs\ML\lib\site-packages\torch\optim\optimizer.py", line 391, in wrapper out = func(*args, **kwargs) File "C:\Users\User1\anaconda3\envs\ML\lib\site-packages\torch\optim\optimizer.py", line 76, in _use_grad ret = func(self, *args, **kwargs) File "C:\Users\User1\anaconda3\envs\ML\lib\site-packages\torch\optim\adam.py", line 168, in step adam( File "C:\Users\User1\anaconda3\envs\ML\lib\site-packages\torch\optim\adam.py", line 318, in adam func(params, File "C:\Users\User1\anaconda3\envs\ML\lib\site-packages\torch\optim\adam.py", line 581, in _multi_tensor_adam exp_avg_sq_sqrt = torch._foreach_sqrt(device_exp_avg_sqs) torch.cuda.OutOfMemoryError: CUDA out of memory. Tried to allocate 2.27 GiB. GPU
解决方案
针对你的情况,以下是几个无需缩小帧尺寸的有效方案:
1. 切换到CNN特征提取+MLP头部,而非纯MLP策略
你当前用的MlpPolicy会把整个视频张量展平成一维向量,产生360230*230=9,522,000维输入,参数和计算量爆炸。换成CnnPolicy或自定义CNN特征提取器,先压缩视频的时空特征再传入MLP:
from stable_baselines3.common.torch_layers import BaseFeaturesExtractor import torch.nn as nn import torch class CustomCNN(BaseFeaturesExtractor): def __init__(self, observation_space: gym.spaces.Box, features_dim: int = 256): super().__init__(observation_space, features_dim) # 输入形状:(3,60,230,230) -> 通道数3,时间步60,高230,宽230 self.cnn = nn.Sequential( # 时空卷积提取帧间+帧内特征 nn.Conv3d(3, 16, kernel_size=(3, 3, 3), stride=(1, 2, 2), padding=(1, 1, 1)), nn.ReLU(), nn.MaxPool3d(kernel_size=(2, 2, 2), stride=(2, 2, 2)), nn.Conv3d(16, 32, kernel_size=(3, 3, 3), stride=(1, 2, 2), padding=(1, 1, 1)), nn.ReLU(), nn.MaxPool3d(kernel_size=(2, 2, 2), stride=(2, 2, 2)), nn.Flatten(), ) # 计算Flatten后的维度 with torch.no_grad(): n_flatten = self.cnn(torch.as_tensor(observation_space.sample()[None]).float()).shape[1] self.linear = nn.Sequential(nn.Linear(n_flatten, features_dim), nn.ReLU()) def forward(self, observations: torch.Tensor) -> torch.Tensor: # 调整维度适配Conv3d输入格式 x = observations.permute(0, 1, 2, 3, 4) return self.linear(self.cnn(x)) # 改用自定义CNN策略 model = PPO("CnnPolicy", env, verbose=1, n_steps=64, batch_size=4, policy_kwargs={"features_extractor_class": CustomCNN}, tensorboard_log="./ppo_tt_tensorboard/")
2. 启用梯度累积
设置gradient_accumulation_steps,把多个小batch的梯度累积后再更新参数,等价于大batch但显存占用更低:
model = PPO("CnnPolicy", env, verbose=1, n_steps=64, batch_size=4, gradient_accumulation_steps=4, tensorboard_log="./ppo_tt_tensorboard/")
此配置每次更新用4个batch(共16样本)的梯度,但显存仅占用1个batch的量。
3. 启用混合精度训练
用FP16半精度训练,PyTorch的torch.cuda.amp可自动混合精度,Stable Baselines3支持直接适配:
from stable_baselines3.common.vec_env import VecTransposeImage from torch.cuda.amp import GradScaler # 转换图像格式为CNN期望的通道优先 env = DummyVecEnv([lambda: TableTennisEnv(frame_tensors, labels)]) env = VecTransposeImage(env) model = PPO("CnnPolicy", env, verbose=1, n_steps=64, batch_size=4, device="cuda", tensorboard_log="./ppo_tt_tensorboard/") scaler = GradScaler() # 自定义训练循环启用混合精度(替代model.learn) def train_with_amp(model, total_timesteps): model.policy.train() for _ in range(total_timesteps // model.n_steps): with torch.cuda.amp.autocast(): rollout = model.collect_rollouts(env, n_rollout_steps=model.n_steps, callback=None) values, log_prob, entropy = model.policy.evaluate_actions(rollout.observations, rollout.actions) advantages = rollout.advantages returns = rollout.returns policy_loss = -torch.mean(log_prob * advantages) value_loss = torch.mean((values - returns) ** 2) loss = policy_loss + 0.5 * value_loss - 0.01 * entropy scaler.scale(loss).backward() scaler.step(model.policy.optimizer) scaler.update() model.policy.optimizer.zero_grad() train_with_amp(model, total_timesteps)
4. 优化数据加载和内存管理
- 不要预先加载所有
frame_tensors到内存,在TableTennisEnv的step和reset方法中按需从磁盘读取当前视频张量。 - 训练循环中定期调用
torch.cuda.empty_cache(),释放无用显存占用。
5. 调整PPO轨迹长度
减小n_steps(比如从64降到32),减少每次收集的轨迹数据量,降低显存占用:
model = PPO("CnnPolicy", env, verbose=1, n_steps=32, batch_size=4, tensorboard_log="./ppo_tt_tensorboard/")
内容的提问来源于stack exchange,提问作者Paarth Jha
相关产品推荐
相关产品推荐

