Add tiling on vae decode for qwen image models when low vram flag is on

This commit is contained in:
Jaret Burkett
2026-07-10 09:07:29 -06:00
parent e7951ad29e
commit fe82487187
3 changed files with 28 additions and 0 deletions

View File

@@ -284,6 +284,11 @@ class QwenImageModel(BaseModel):
sc = self.get_bucket_divisibility()
gen_config.width = int(gen_config.width // sc * sc)
gen_config.height = int(gen_config.height // sc * sc)
if self.model_config.low_vram:
# set vae to tile decode
pipeline.vae.enable_tiling()
img = pipeline(
prompt_embeds=conditional_embeds.text_embeds,
prompt_embeds_mask=conditional_embeds.attention_mask.to(
@@ -302,6 +307,11 @@ class QwenImageModel(BaseModel):
callback_on_step_end=callback_on_step_end,
**extra,
).images[0]
if self.model_config.low_vram:
# restore no tiling
pipeline.vae.disable_tiling()
return img
def get_noise_prediction(

View File

@@ -115,6 +115,10 @@ class QwenImageEditModel(QwenImageModel):
return {"latents": latents}
if self.model_config.low_vram:
# set vae to tile decode
pipeline.vae.enable_tiling()
img = pipeline(
image=control_img,
prompt_embeds=conditional_embeds.text_embeds,
@@ -134,6 +138,11 @@ class QwenImageEditModel(QwenImageModel):
callback_on_step_end=callback_on_step_end,
**extra,
).images[0]
if self.model_config.low_vram:
# restore no tiling
pipeline.vae.disable_tiling()
return img
def condition_noisy_latents(

View File

@@ -133,6 +133,10 @@ class QwenImageEditPlusModel(QwenImageModel):
return {"latents": latents}
if self.model_config.low_vram:
# set vae to tile decode
pipeline.vae.enable_tiling()
img = pipeline(
image=control_img_list,
prompt_embeds=conditional_embeds.text_embeds,
@@ -153,6 +157,11 @@ class QwenImageEditPlusModel(QwenImageModel):
do_cfg_norm=gen_config.do_cfg_norm,
**extra,
).images[0]
if self.model_config.low_vram:
# restore no tiling
pipeline.vae.disable_tiling()
return img
def condition_noisy_latents(