集成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

