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

Stable Diffusion代码执行pipe.to(device)时卡住,寻求解决办法

解决4GB显存下Stable Diffusion加载卡住的问题

核心原因

4GB显存不足以一次性加载完整的Stable Diffusion v1-4 fp16模型,导致pipe.to(device)执行时显存耗尽,进程卡住无响应。

具体解决方案

1. 启用模型分片与CPU卸载

修改模型加载后的处理逻辑,去掉直接的pipe.to(device),改用diffusers提供的显存优化工具:

# 初始化pipe时关闭安全检查器以节省显存
pipe = StableDiffusionPipeline.from_pretrained(
    model_id,
    revision="fp16",
    torch_dtype=torch.float16,
    use_auth_token=auth_token,
    safety_checker=None,
    requires_safety_checker=False
)

# 启用顺序CPU卸载,自动在需要时将模型部分加载到GPU,用完放回CPU
pipe.enable_sequential_cpu_offload()

# 启用注意力切片,进一步降低显存占用
pipe.unet.enable_attention_slicing()

2. 安装并启用xFormers优化

xFormers能大幅优化注意力层的显存使用,适合小显存设备:

  • 安装xFormers(以CUDA 11.7为例,根据自身CUDA版本调整):
    pip install xformers==0.0.20
    
  • 代码中启用优化:
    pipe.enable_xformers_memory_efficient_attention()
    

3. 清理显存占用

确保没有其他进程占用GPU显存:

  • 执行nvidia-smi命令查看显存使用情况
  • 关闭未使用的PyTorch进程或其他GPU应用,释放显存

4. 改用更轻量化的模型版本

如果上述优化仍无法解决,可尝试显存占用更低的模型:

  • 将model_id改为runwayml/stable-diffusion-v1-5(优化更好,显存占用略低)
  • 或使用stabilityai/stable-diffusion-2-1-base(基础版模型体积更小)

修改后的完整代码示例

from auth_token import auth_token
from fastapi import FastAPI, Response
from fastapi.middleware.cors import CORSMiddleware
import torch
from torch import autocast
from diffusers import StableDiffusionPipeline
from io import BytesIO
import base64

from torch.cuda import empty_cache

app = FastAPI()

app.add_middleware(
    CORSMiddleware,
    allow_credentials=True,
    allow_origins=["*"],
    allow_methods=["*"],
    allow_headers=["*"]
)

device = "cuda"
model_id = "runwayml/stable-diffusion-v1-5"  # 改用优化后的v1-5模型
pipe = StableDiffusionPipeline.from_pretrained(
    model_id,
    revision="fp16",
    torch_dtype=torch.float16,
    use_auth_token=auth_token,
    safety_checker=None,
    requires_safety_checker=False
)

# 启用显存优化
pipe.enable_sequential_cpu_offload()
pipe.unet.enable_attention_slicing()
# 若已安装xFormers,取消注释以下行
# pipe.enable_xformers_memory_efficient_attention()

@app.get("/")
def generate(prompt: str):
    with autocast(device):
        image = pipe(prompt, guidance_scale=7.5).images[0]  # 适当降低guidance_scale减少显存压力

    image.save("testimage.png")
    empty_cache()
 
    return {"out": "hello World"}

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 22:15:22