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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 03:07:05