Stable Cascade模型生图耗时超7小时,求M2 Pro环境下提速方案
我正在使用Stable Cascade模型,初始代码如下:
from diffusers import StableCascadeCombinedPipeline print("LOADING MODEL") pipe = StableCascadeCombinedPipeline.from_pretrained("stabilityai/stable-cascade", variant="bf16", torch_dtype=torch.bfloat16) print("MODEL LOADED") prompt = "a lawyer" pipe( prompt=prompt, negative_prompt="", num_inference_steps=10, prior_num_inference_steps=20, prior_guidance_scale=3.0, width=1024, height=1024, ).images[0].save("cascade-combined2.png")
模型加载几乎瞬间完成,但生图耗时超过7小时,运行日志:
Loading pipeline components...: 100%|█████████████████████████████████████████| 5/5 [00:00<00:00, 9.73it/s] Loading pipeline components...: 100%|█████████████████████████████████████████| 6/6 [00:00<00:00, 11.08it/s] MODEL LOADED 0%| | 0/20 [00:00<?, ?it/s] 0%| | 0/20 [04:10<?, ?it/s]
我的运行环境是Apple M2 Pro(32 GB)、Python 3.10.2,需要生成约50张图,当前速度无法满足需求,求提速方法。
补充修改后的代码:
pipe = StableCascadeCombinedPipeline.from_pretrained("stabilityai/stable-cascade", variant="bf16", torch_dtype=torch.float32) device = torch.device('mps') pipe.to(device) prompt = "a football" pipe( prompt=prompt, negative_prompt="", num_inference_steps=10, prior_num_inference_steps=20, prior_guidance_scale=3.0, width=1024, height=1024, ).images[0].save("cascade-combined2.png")
针对Apple M系列芯片的优化建议如下:
优化MPS设备设置
确保PyTorch版本在2.0及以上,M2 Pro对新版PyTorch的MPS支持更完善。代码中添加MPS专属优化:import torch torch.backends.mps.enabled = True torch.backends.mps.allow_infconv = True所有模型组件和张量都移至MPS设备,避免CPU与MPS间频繁数据传输。
调整模型精度与变体
M2 Pro对bf16精度支持更高效,建议改回torch.bfloat16并配合模型变体加载,同时关闭注意力切片以提升速度(需确保内存充足):pipe = StableCascadeCombinedPipeline.from_pretrained( "stabilityai/stable-cascade", variant="bf16", torch_dtype=torch.bfloat16 ) pipe.to("mps") pipe.enable_attention_slicing(None) # 关闭注意力切片减少推理步数与分辨率
降低prior_num_inference_steps(如从20降至10)和num_inference_steps(如从10降至5),对生成效果影响微小但能大幅提速;若无需1024x1024分辨率,降到768x768可提升约50%速度。批量生成图片
一次性传入多个prompt进行批量处理,减少模型重复初始化开销:prompts = ["a lawyer", "a football"] + ["你的其他prompt"] * 48 results = pipe( prompt=prompts, negative_prompt=[""]*len(prompts), num_inference_steps=5, prior_num_inference_steps=10, width=768, height=768 ) for i, img in enumerate(results.images): img.save(f"cascade-combined_{i}.png")更新依赖库版本
升级diffusers到0.24.0以上版本,官方针对Stable Cascade和MPS做了专项优化;同时更新accelerate、torch、transformers等依赖:pip install --upgrade diffusers accelerate torch transformers safetensors内存优化
启用模型CPU卸载功能,将暂时不用的模型组件移至CPU,释放MPS内存避免因内存不足降速:pipe.enable_model_cpu_offload()关闭后台不必要程序,让算力集中在生图任务上。
内容的提问来源于stack exchange,提问作者x89

