如何优化MarioGPT脚本以解决关卡生成卡顿问题
解决MarioGPT关卡生成速度过慢的问题
核心问题分析
你当前的代码每秒仅生成10步,主要瓶颈大概率在模型推理设备和采样配置上,和temperature、prompt内容关联不大,重新训练模型反而会增加额外开销,完全没必要。
具体优化方案
强制启用GPU加速
MarioLM默认可能未自动调用GPU,初始化时明确指定设备:mario_lm = MarioLM(device="cuda") # 适用于NVIDIA GPU # 苹果芯片用户替换为: mario_lm = MarioLM(device="mps")验证是否生效:打印
print(mario_lm.device),确认输出为cuda或mps。调整采样的步数与批次配置
单次生成1400步的计算压力较大,可拆分生成批次,同时根据显存情况合理设置batch_size:mario_lm = MarioLM(device="cuda") def my_function(): # 拆分生成任务,每次生成200步,分7次完成 result = mario_lm.sample( prompts=["some pipes, some enemies, some blocks, low elevation"], num_steps=200, batch_size=2 # 显存充足可适当增大,提升生成效率 ) for _ in range(6): result = mario_lm.sample( prompts=[result.level_str], num_steps=200, batch_size=2 ) result.play() my_function()启用半精度推理
开启float16半精度计算,能大幅提升推理速度且几乎不影响生成质量:mario_lm = MarioLM(device="cuda", dtype=torch.float16)替换递归调用为循环
你当前用递归调用my_function(),会持续创建函数栈,长期运行可能引发内存泄漏拖慢速度,改成循环更高效:mario_lm = MarioLM(device="cuda", dtype=torch.float16) while True: mario_lm.sample( prompts=["some pipes, some enemies, some blocks, low elevation"], num_steps=1400, batch_size=1 ).play()
额外排查点
- 检查GPU显存占用:如果显存不足会自动 fallback 到CPU,导致速度暴跌。用
nvidia-smi(Windows用户用任务管理器)查看显存使用情况,关闭其他占用显存的程序。 - 升级MarioGPT版本:旧版本可能存在推理效率缺陷,执行
pip install --upgrade mario-gpt更新到最新版。
内容的提问来源于stack exchange,提问作者user22155681
相关产品推荐
相关产品推荐

