@@ -1491,6 +1491,8 @@ def __init__(
14911491 freeze_encoder : bool ,
14921492 cpu_offload : bool = False ,
14931493 use_tiling : bool = False ,
1494+ tile_size : Union [int , Tuple [int , int ]] = 512 ,
1495+ tile_overlap : Union [int , Tuple [int , int ]] = 64 ,
14941496 ** kwargs ,
14951497 ):
14961498 self .spatial_downsample_factor = spatial_downsample_factor
@@ -1499,6 +1501,14 @@ def __init__(
14991501 self .cpu_offload = cpu_offload
15001502 super ().__init__ (* args , ** kwargs )
15011503 self .use_tiling = use_tiling
1504+ self .tile_size = self ._as_pair (tile_size )
1505+ self .tile_overlap = self ._as_pair (tile_overlap )
1506+
1507+ @staticmethod
1508+ def _as_pair (value : Union [int , Tuple [int , int ]]) -> Tuple [int , int ]:
1509+ if isinstance (value , tuple ):
1510+ return int (value [0 ]), int (value [1 ])
1511+ return int (value ), int (value )
15021512
15031513 def forward (self , x : torch .FloatTensor ) -> CausalAutoencoderOutput :
15041514 with torch .no_grad () if self .freeze_encoder else nullcontext ():
@@ -1572,10 +1582,10 @@ def vae_encode(self, samples: List[Tensor]) -> List[Tensor]:
15721582 if hasattr (self , "preprocess" ):
15731583 sample = self .preprocess (sample )
15741584 if use_sample :
1575- latent = self .encode (sample , tiled = self .use_tiling ).latent
1585+ latent = self .encode (sample , tiled = self .use_tiling , tile_size = self . tile_size , tile_overlap = self . tile_overlap ).latent
15761586 else :
15771587 # Deterministic vae encode, only used for i2v inference (optionally)
1578- latent = self .encode (sample , tiled = self .use_tiling ).posterior .mode ().squeeze (2 )
1588+ latent = self .encode (sample , tiled = self .use_tiling , tile_size = self . tile_size , tile_overlap = self . tile_overlap ).posterior .mode ().squeeze (2 )
15791589 latent = latent .unsqueeze (2 ) if latent .ndim == 4 else latent
15801590 latent = rearrange (latent , "b c ... -> b ... c" )
15811591 latent = (latent - shift ) * scale
@@ -1614,7 +1624,7 @@ def vae_decode(self, latents: List[Tensor]) -> List[Tensor]:
16141624 latent = latent / scale + shift
16151625 latent = rearrange (latent , "b ... c -> b c ..." )
16161626 latent = latent .squeeze (2 )
1617- sample = self .decode (latent , tiled = self .use_tiling ).sample
1627+ sample = self .decode (latent , tiled = self .use_tiling , tile_size = self . tile_size , tile_overlap = self . tile_overlap ).sample
16181628 if hasattr (self , "postprocess" ):
16191629 sample = self .postprocess (sample )
16201630 samples .append (sample )
@@ -1642,6 +1652,8 @@ def attn_video_vae_v3_s8_c16_t4_inflation_sd3_init(
16421652 strict : bool = True ,
16431653 cpu_offload : bool = False ,
16441654 use_tiling : bool = False ,
1655+ tile_size : Union [int , Tuple [int , int ]] = 512 ,
1656+ tile_overlap : Union [int , Tuple [int , int ]] = 64 ,
16451657) -> VideoAutoencoderKLWrapper :
16461658 """Example: initialize VideoAutoencoderKLWrapper with SD3 inflation config params."""
16471659 model = VideoAutoencoderKLWrapper (
@@ -1673,6 +1685,8 @@ def attn_video_vae_v3_s8_c16_t4_inflation_sd3_init(
16731685 freeze_encoder = False ,
16741686 cpu_offload = cpu_offload ,
16751687 use_tiling = use_tiling ,
1688+ tile_size = tile_size ,
1689+ tile_overlap = tile_overlap ,
16761690 )
16771691
16781692 if weights_path is not None :
0 commit comments