Google Colab运行SDXL Base+Refiner时会话崩溃求助
在Google Colab中同时运行SDXL Base+Refiner避免会话崩溃的解决方案
问题描述
我使用Hugging Face提供的示例Python代码尝试同时运行SDXL Base和Refiner模型,但每次图像即将生成时Google Colab会话都会崩溃。单独运行Base模型是成功的,希望找到能同时运行两者的方法。
示例代码:
from diffusers import DiffusionPipeline import torch # load both base & refiner base = DiffusionPipeline.from_pretrained( "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16, variant="fp16", use_safetensors=True ) base.to("cuda") refiner = DiffusionPipeline.from_pretrained( "stabilityai/stable-diffusion-xl-refiner-1.0", text_encoder_2=base.text_encoder_2, vae=base.vae, torch_dtype=torch.float16, use_safetensors=True, variant="fp16", ) refiner.to("cuda") # Define how many steps and what % of steps to be run on each experts (80/20) here n_steps = 40 high_noise_frac = 0.8 prompt = "A majestic lion jumping from a big stone at night" # run both experts image = base( prompt=prompt, num_inference_steps=n_steps, denoising_end=high_noise_frac, output_type="latent", ).images image = refiner( prompt=prompt, num_inference_steps=n_steps, denoising_start=high_noise_frac, image=image, ).images[0] image
可行解决方案
- 启用Colab高RAM模式:点击顶部菜单栏「Runtime」→「Change runtime type」,勾选「High-RAM」选项,获取更多内存,缓解显存/内存不足问题。
- 加载模型时启用自动设备映射:添加
device_map="auto"参数,让模型自动分配设备,减少显存占用。修改后的加载代码如下:
base = DiffusionPipeline.from_pretrained( "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16, variant="fp16", use_safetensors=True, device_map="auto" # 新增参数 ) refiner = DiffusionPipeline.from_pretrained( "stabilityai/stable-diffusion-xl-refiner-1.0", text_encoder_2=base.text_encoder_2, vae=base.vae, torch_dtype=torch.float16, use_safetensors=True, variant="fp16", device_map="auto" # 新增参数 )
- 分阶段卸载模型:生成Base的latent后,先卸载Base模型释放显存,再加载Refiner。示例代码:
# 运行Base模型生成latent image = base( prompt=prompt, num_inference_steps=n_steps, denoising_end=high_noise_frac, output_type="latent", ).images # 卸载Base模型释放显存 del base torch.cuda.empty_cache() # 加载并运行Refiner模型 refiner = DiffusionPipeline.from_pretrained( "stabilityai/stable-diffusion-xl-refiner-1.0", torch_dtype=torch.float16, use_safetensors=True, variant="fp16", device_map="auto" ) image = refiner( prompt=prompt, num_inference_steps=n_steps, denoising_start=high_noise_frac, image=image, ).images[0] image
内容的提问来源于stack exchange,提问作者James Arnold
相关产品推荐
相关产品推荐

