Allow for vae tiling onle without low vram on wan models with a model kwarg
This commit is contained in:
@@ -555,6 +555,11 @@ class Wan2214bModel(Wan21):
|
||||
):
|
||||
# reactivate progress bar since this is slooooow
|
||||
pipeline.set_progress_bar_config(disable=False)
|
||||
|
||||
if self.use_vae_tiling:
|
||||
# set vae to tile decode
|
||||
pipeline.vae.enable_tiling()
|
||||
|
||||
# todo, figure out how to do video
|
||||
output = pipeline(
|
||||
prompt_embeds=conditional_embeds.text_embeds.to(
|
||||
@@ -573,6 +578,10 @@ class Wan2214bModel(Wan21):
|
||||
**extra
|
||||
)[0]
|
||||
|
||||
if self.use_vae_tiling:
|
||||
# restore no tiling
|
||||
pipeline.vae.disable_tiling()
|
||||
|
||||
# shape = [1, frames, channels, height, width]
|
||||
batch_item = output[0] # list of pil images
|
||||
if gen_config.num_frames > 1:
|
||||
|
||||
@@ -210,7 +210,7 @@ class Wan225bModel(Wan21):
|
||||
latent_model_input=latents, first_frame=first_frame_n1p1, vae=self.vae
|
||||
)
|
||||
|
||||
if self.model_config.low_vram:
|
||||
if self.use_vae_tiling:
|
||||
# set vae to tile decode
|
||||
pipeline.vae.enable_tiling()
|
||||
|
||||
@@ -234,7 +234,7 @@ class Wan225bModel(Wan21):
|
||||
**extra,
|
||||
)[0]
|
||||
|
||||
if self.model_config.low_vram:
|
||||
if self.use_vae_tiling:
|
||||
# restore no tiling
|
||||
pipeline.vae.disable_tiling()
|
||||
|
||||
|
||||
@@ -520,6 +520,13 @@ class Wan21(BaseModel):
|
||||
|
||||
return pipeline
|
||||
|
||||
@property
|
||||
def use_vae_tiling(self):
|
||||
# tile the vae decode when sampling if in low vram or explicitly enabled
|
||||
return self.model_config.low_vram or self.model_config.model_kwargs.get(
|
||||
"vae_tiling", False
|
||||
)
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: WanPipeline,
|
||||
@@ -532,6 +539,11 @@ class Wan21(BaseModel):
|
||||
# reactivate progress bar since this is slooooow
|
||||
pipeline.set_progress_bar_config(disable=False)
|
||||
pipeline = pipeline.to(self.device_torch)
|
||||
|
||||
if self.use_vae_tiling:
|
||||
# set vae to tile decode
|
||||
pipeline.vae.enable_tiling()
|
||||
|
||||
# todo, figure out how to do video
|
||||
output = pipeline(
|
||||
prompt_embeds=conditional_embeds.text_embeds.to(
|
||||
@@ -550,6 +562,10 @@ class Wan21(BaseModel):
|
||||
**extra
|
||||
)[0]
|
||||
|
||||
if self.use_vae_tiling:
|
||||
# restore no tiling
|
||||
pipeline.vae.disable_tiling()
|
||||
|
||||
# shape = [1, frames, channels, height, width]
|
||||
batch_item = output[0] # list of pil images
|
||||
if gen_config.num_frames > 1:
|
||||
|
||||
Reference in New Issue
Block a user