Skip to content

Commit f32c5db

Browse files
committed
feat: expose seedvr vae tuning
1 parent 548f068 commit f32c5db

2 files changed

Lines changed: 38 additions & 8 deletions

File tree

lightx2v/models/runners/seedvr/seedvr_runner.py

Lines changed: 21 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -305,6 +305,9 @@ def load_image_encoder(self):
305305
return None
306306

307307
def load_vae_encoder(self):
308+
vae_causal_slice_size = int(self.config.get("vae_causal_slice_size", 4))
309+
vae_memory_limit_gb = float(self.config.get("vae_memory_limit_gb", 0.5))
310+
vae_memory_limit = None if vae_memory_limit_gb <= 0 else vae_memory_limit_gb
308311
vae = attn_video_vae_v3_s8_c16_t4_inflation_sd3_init(
309312
device=AI_DEVICE,
310313
dtype=GET_DTYPE(),
@@ -314,10 +317,18 @@ def load_vae_encoder(self):
314317
strict=False,
315318
cpu_offload=self.config.get("cpu_offload", False),
316319
use_tiling=self.config.get("use_tiling_vae", False),
320+
tile_size=int(self.config.get("vae_tile_size", 512)),
321+
tile_overlap=int(self.config.get("vae_tile_overlap", 64)),
317322
)
318323
vae.requires_grad_(False).eval()
319-
vae.set_causal_slicing(split_size=4, memory_device="same")
320-
vae.set_memory_limit(conv_max_mem=0.5, norm_max_mem=0.5)
324+
vae.set_causal_slicing(split_size=vae_causal_slice_size if vae_causal_slice_size > 0 else None, memory_device="same")
325+
vae.set_memory_limit(conv_max_mem=vae_memory_limit, norm_max_mem=vae_memory_limit)
326+
logger.info(
327+
f"[SeedVRRunner] VAE config: tiling={self.config.get('use_tiling_vae', False)}, "
328+
f"tile={self.config.get('vae_tile_size', 512)}, overlap={self.config.get('vae_tile_overlap', 64)}, "
329+
f"causal_slice={vae_causal_slice_size if vae_causal_slice_size > 0 else 'off'}, "
330+
f"memory_limit={vae_memory_limit_gb if vae_memory_limit_gb > 0 else 'off'}GiB"
331+
)
321332
return vae
322333

323334
def load_vae_decoder(self):
@@ -340,9 +351,14 @@ def run_vae_decoder(self, latents):
340351
if self._ori_length < sample.shape[0]:
341352
sample = sample[: self._ori_length]
342353

343-
# color fix
344-
input = rearrange(self._input[:, None], "c t h w -> t c h w") if self._input.ndim == 3 else rearrange(self._input, "c t h w -> t c h w")
345-
sample = wavelet_reconstruction(sample.to("cpu"), input[: sample.size(0)].to("cpu"))
354+
color_fix = str(self.config.get("color_fix", "cpu")).lower()
355+
if color_fix not in ("cpu", "gpu", "off"):
356+
logger.warning(f"[SeedVRRunner] Unknown color_fix={color_fix}; fallback to cpu")
357+
color_fix = "cpu"
358+
if color_fix != "off":
359+
input = rearrange(self._input[:, None], "c t h w -> t c h w") if self._input.ndim == 3 else rearrange(self._input, "c t h w -> t c h w")
360+
fix_device = torch.device("cpu") if color_fix == "cpu" else sample.device
361+
sample = wavelet_reconstruction(sample.to(fix_device), input[: sample.size(0)].to(fix_device))
346362
sample = rearrange(sample[:, None], "t c h w -> c t h w") if sample.ndim == 3 else rearrange(sample, "t c h w -> c t h w")
347363
sample = sample[None, :]
348364

lightx2v/models/video_encoders/hf/seedvr/attn_video_vae.py

Lines changed: 17 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)