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

集成Normalizing flow的模型基强化学习代码训练过慢求助

模型基强化学习训练速度异常缓慢的排查方向

我是模型基强化学习(Model-based reinforcement learning)研究者,为通过流网络获取更优模拟样本,在代码中加入了Normalizing flow模型。代码集成调试后可正常运行,但训练速度极慢:从昨晚10点到今早10点仅完成3个epoch(3000步)。已尝试调整流网络batch_size(100→256→500)、减少文件IO、更换Tesla-v100-SXM2-32GB显卡,但收效甚微。以下是train函数代码:

def train(args, env_sampler, predict_env, agent, env_pool, model_pool, couple_flow, rewarder, model_agent, rollout_schedule):
    total_step = 0
    rollout_length =1
    rollout_depth =2
    exploration_before_start(args, env_sampler, env_pool, agent)

    for epoch_step in range(args.num_epoch):
        start_step = total_step
        train_policy_steps = 0
        for i in count():
            cur_step = total_step - start_step

            if cur_step >= args.epoch_length and len(env_pool) > args.min_pool_size:  # epoch_length is 1000
                break

            if args.use_algo == 'discriminator':
                if i == 0:
                    train_predict_model(args, env_pool, predict_env)
                    print("completed model!!!!!!!!!!!!")
                #function2
                train_couple_flow(args, env_pool, predict_env, agent, rollout_depth, rewarder, total_step)
                if cur_step > 0 and cur_step % args.model_train_freq == 0 and args.real_ratio < 1.0:
                    print("start train model")
                    train_predict_model(args, env_pool, predict_env)
                    print("end train model")
                    new_rollout_length = set_rollout_length(epoch_step, rollout_schedule)
                    if rollout_length != new_rollout_length:
                        rollout_length = new_rollout_length
                        model_pool = resize_model_pool(args, rollout_length, model_pool)
                    print("start rollouting")

                    MPC_rollout_model(args, predict_env, agent, model_pool, env_pool, rollout_length, rewarder, total_step)
                    # rollout_model(args, predict_env, agent, model_pool, env_pool, rollout_length)
                    print("end rollouting")
            elif args.use_algo == 'flowrl':
                if i == 0:
                    train_predict_model(args, env_pool, predict_env)
                #function1
                train_predict_model_by_couple_flow(args, env_pool, predict_env, agent, rollout_depth, rewarder, model_agent, cur_step, total_step)
                if cur_step > 0 and cur_step % args.model_train_freq == 0 and args.real_ratio < 1.0:
                    #  if cur_step > 0 and cur_step % 500 == 0 and args.real_ratio < 1.0:

                    new_rollout_length = set_rollout_length(epoch_step, rollout_schedule)
                    if rollout_length != new_rollout_length:
                        rollout_length = new_rollout_length
                        model_pool = resize_model_pool(args, rollout_length, model_pool)

                    rollout_model(args, predict_env, agent, model_pool, env_pool, rollout_length)
            elif args.use_algo == 'mbpo':
                if cur_step > 0 and cur_step % args.model_train_freq == 0 and args.real_ratio < 1.0:
                    train_predict_model(args, env_pool, predict_env)

                    new_rollout_length = set_rollout_length(epoch_step, rollout_schedule)
                    if rollout_length != new_rollout_length:
                        rollout_length = new_rollout_length
                        model_pool = resize_model_pool(args, rollout_length, model_pool)

                    rollout_model(args, predict_env, agent, model_pool, env_pool, rollout_length)


            cur_state, action, next_state, reward, done, info = env_sampler.sample(agent)
            env_pool.push(cur_state, action, reward, next_state, done)

            if len(env_pool) > args.min_pool_size:
                train_policy_steps += train_policy_repeats(args, total_step, train_policy_steps, cur_step, env_pool, model_pool, agent)

            total_step += 1

            if total_step % args.epoch_length == 0:
                '''
                avg_reward_len = min(len(env_sampler.path_rewards), 5)
                avg_reward = sum(env_sampler.path_rewards[-avg_reward_len:]) / avg_reward_len
                logging.info("Step Reward: " + str(total_step) + " " + str(env_sampler.path_rewards[-1]) + " " + str(avg_reward))
                print(total_step, env_sampler.path_rewards[-1], avg_reward)
                '''
                env_sampler.current_state = None
                sum_reward = 0
                done = False
                test_step = 0

                while (not done) and (test_step != args.max_path_length):
                    cur_state, action, next_state, reward, done, info = env_sampler.sample(agent, eval_t=True)
                    sum_reward += reward
                    test_step += 1
                # logger.record_tabular("total_step", total_step)
                # logger.record_tabular("sum_reward", sum_reward)
                # logger.dump_tabular()

                folder_path = f"./results/{args.env_name}"
                if not os.path.exists(folder_path):
                    os.makedirs(folder_path)

                if args.use_algo == 'discriminator':
                    file_name = f"{folder_path}/{args.env_name}_discriminator_{now02}.txt"
                elif args.use_algo == 'flowrl':
                    file_name = f"{folder_path}/{args.env_name}_flowRL_{now02}.txt"
                elif args.use_algo == 'mbpo':
                    file_name = f"{folder_path}/{args.env_name}_mbpo_{now02}.txt"

                with open(file_name, "a") as file:
                    file.write(f"{total_step}\t{sum_reward}\n")

                logging.info("Step Reward: " + str(total_step) + " " + str(sum_reward))
                print(total_step, sum_reward)
                torch.cuda.empty_cache()

可能的原因排查

1. 流网络训练频率过高

从代码逻辑看,discriminator和flowrl模式下,每一步循环都会调用流相关训练函数(train_couple_flow或train_predict_model_by_couple_flow)。Normalizing Flow本身是计算密集型模型,每步都训练会直接把计算资源占满。对比MBPO模式仅在model_train_freq间隔训练模型,流模式的训练频率显然不合理,这是核心瓶颈之一。

2. 流网络内部训练逻辑低效

即使调整了batch_size,流网络本身的实现可能存在优化空间:

  • 未启用混合精度训练:流网络的对数密度计算涉及大量矩阵运算,开启torch.cuda.amp混合精度能大幅降低计算开销和显存占用。
  • 流结构过于复杂:如果使用了过多的Flow变换层(比如深层RealNVP、Glow结构),单步训练的计算量会呈指数级上升。
  • 数据传输冗余:env_pool的数据读取如果频繁在CPU和GPU之间切换,每次训练都要做数据拷贝,会产生大量延迟。

3. Rollout过程计算冗余

MPC_rollout_model或rollout_model可能存在以下问题:

  • Rollout规模过大:如果每次rollout生成上万条模拟轨迹,每条轨迹走几十步,计算量会直接爆炸。
  • MPC规划成本高:discriminator模式下的MPC rollout,若每步都用迭代次数多的算法(如CEM)做规划,会占用大量计算资源。
  • 模型池频繁重建:如果rollout_length频繁变化,resize_model_pool会反复进行内存分配和数据拷贝,拖慢训练节奏。

4. 策略训练次数失控

train_policy_repeats函数的逻辑未知,但如果内部每次调用都会进行多轮策略更新(比如PPO的多轮epoch),且没有合理的终止条件,会导致每步循环内策略训练的计算量过载。

5. CPU-GPU同步瓶颈

代码中存在大量print和logging操作,这些操作会触发GPU到CPU的同步(必须等待GPU计算完成才能打印状态),频繁的同步会打断GPU的异步计算,导致训练停滞。另外,虽然torch.cuda.empty_cache()在epoch结束调用合理,但如果内部函数存在频繁缓存清理,也会影响效率。

6. 环境采样开销

如果env_sampler.sample(agent)是在CPU环境中运行(比如MuJoCo、Atari的CPU版本),每步采样的延迟会叠加,尤其是环境本身计算复杂时,会拖慢整个训练循环。可以尝试批量采样多个环境步,或把环境计算移到GPU(如果支持)。


验证建议

  • 先注释掉流网络的训练函数,观察训练速度是否恢复,确认流网络是核心瓶颈。
  • 给流网络训练添加频率控制,比如每10步或50步训练一次,而非每步都训。
  • 用torch.profiler或cProfile分析代码,定位耗时最长的函数(重点看流训练、rollout、策略训练部分)。
  • 检查流网络实现,确保所有操作都在GPU上进行,避免不必要的CPU-GPU数据传输。

内容的提问来源于stack exchange,提问作者user23997403

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 11:02:05