From d4380e441c9e7a936029ddde8b2cc855b6149acc Mon Sep 17 00:00:00 2001 From: Super User Date: Sat, 8 Aug 2026 22:48:44 +0800 Subject: [PATCH 1/3] feat(flux2): optimize distributed encoder and decoder --- lightx2v/models/runners/flux2/flux2_runner.py | 96 ++++++++++++++++--- .../models/video_encoders/hf/flux2/vae.py | 5 +- 2 files changed, 88 insertions(+), 13 deletions(-) diff --git a/lightx2v/models/runners/flux2/flux2_runner.py b/lightx2v/models/runners/flux2/flux2_runner.py index 8d63b8cac..7205df684 100644 --- a/lightx2v/models/runners/flux2/flux2_runner.py +++ b/lightx2v/models/runners/flux2/flux2_runner.py @@ -4,6 +4,7 @@ import numpy as np import torch +import torch.distributed as dist from loguru import logger from lightx2v.models.networks.flux2.model import Flux2DevTransformerModel, Flux2KleinTransformerModel @@ -11,6 +12,7 @@ from lightx2v.models.schedulers.flux2.feature_caching.scheduler import Flux2DevSchedulerCaching, Flux2SchedulerCaching from lightx2v.models.schedulers.flux2.scheduler import Flux2DevScheduler, Flux2Scheduler from lightx2v.models.video_encoders.hf.flux2.vae import Flux2VAE +from lightx2v.utils.envs import GET_DTYPE from lightx2v.utils.profiler import ProfilingContext4DebugL1, ProfilingContext4DebugL2 from lightx2v.utils.registry_factory import RUNNER_REGISTER from lightx2v.utils.utils import is_main_process @@ -48,10 +50,24 @@ def _get_scheduler_class(self): @ProfilingContext4DebugL2("Load models") def load_model(self): - self.text_encoders = self.load_text_encoder() + self.text_encoders = self._load_text_encoder_on_this_rank() self.vae = self.load_vae() + if self._rank0_text_encoder_broadcast_enabled(): + torch_device_module.synchronize() + if AI_DEVICE == "cuda": + dist.barrier(device_ids=[torch.cuda.current_device()]) + else: + dist.barrier() self.model = self.load_transformer() + def _rank0_text_encoder_broadcast_enabled(self): + return False + + def _load_text_encoder_on_this_rank(self): + if self._rank0_text_encoder_broadcast_enabled() and dist.get_rank() != 0: + return None + return self.load_text_encoder() + def load_vae(self): return Flux2VAE(self.config) @@ -75,9 +91,9 @@ def init_modules(self): def _run_input_encoder_local_t2i(self): prompt = self.input_info.prompt if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): - self.text_encoders = self.load_text_encoder() + self.text_encoders = self._load_text_encoder_on_this_rank() text_encoder_output = self.run_text_encoder(prompt, neg_prompt=self.input_info.negative_prompt) - if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): + if (self.config.get("lazy_load", False) or self.config.get("unload_modules", False)) and self.text_encoders is not None: del self.text_encoders[0] torch_device_module.empty_cache() gc.collect() @@ -90,9 +106,9 @@ def _run_input_encoder_local_t2i(self): def _run_input_encoder_local_i2i(self): prompt = self.input_info.prompt if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): - self.text_encoders = self.load_text_encoder() + self.text_encoders = self._load_text_encoder_on_this_rank() text_encoder_output = self.run_text_encoder(prompt, neg_prompt=self.input_info.negative_prompt) - if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): + if (self.config.get("lazy_load", False) or self.config.get("unload_modules", False)) and self.text_encoders is not None: del self.text_encoders[0] image_path = self.input_info.image_path @@ -338,7 +354,11 @@ def run_vae_decoder(self, latents): if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): self.vae = self.load_vae() - images = self._decode_latents_with_vae(latents) + decode_rank = self._distributed_vae_decode_rank() + if decode_rank is None: + images = self._decode_latents_with_vae(latents) + else: + images = self._decode_vae_on_rank(latents, decode_rank) if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): del self.vae @@ -347,7 +367,35 @@ def run_vae_decoder(self, latents): return images - def _decode_latents_with_vae(self, latents): + def _distributed_vae_decode_rank(self): + if not dist.is_initialized() or dist.get_world_size() == 1: + return None + return self.config.get("vae_decode_rank") + + def _decode_vae_on_rank(self, latents, decode_rank): + torch_device_module.synchronize() + torch_device_module.empty_cache() + + rank = dist.get_rank() + if rank == decode_rank: + logger.info(f"[Flux2] Decoding VAE on global rank {decode_rank} and broadcasting images") + raw_images = self._decode_latents_with_vae(latents, output_type="latent").contiguous() + image_shape = torch.tensor(raw_images.shape, dtype=torch.long, device=raw_images.device) + else: + image_shape = torch.empty(4, dtype=torch.long, device=latents.device) + + dist.broadcast(image_shape, src=decode_rank, group=dist.group.WORLD) + if rank != decode_rank: + raw_images = torch.empty(tuple(image_shape.tolist()), dtype=GET_DTYPE(), device=latents.device) + dist.broadcast(raw_images, src=decode_rank, group=dist.group.WORLD) + + if self.input_info.return_result_tensor: + return self.vae.image_processor.postprocess(raw_images, output_type="pt") + if is_main_process(): + return self.vae.image_processor.postprocess(raw_images, output_type="pil") + return None + + def _decode_latents_with_vae(self, latents, output_type=None): B, _, C = latents.shape H = int((self.input_info.latent_image_ids[0, :, 1].max() + 1).item()) @@ -363,7 +411,7 @@ def _decode_latents_with_vae(self, latents): latents = latents.permute(0, 1, 4, 2, 5, 3) latents = latents.reshape(B, C // 4, H * 2, W * 2) - return self.vae.decode(latents, self.input_info) + return self.vae.decode(latents, self.input_info, output_type=output_type) @ProfilingContext4DebugL1("RUN pipeline") def run_pipeline(self, input_info): @@ -436,6 +484,9 @@ def run_text_encoder(self, text, image_list=None, neg_prompt=None): @RUNNER_REGISTER("flux2_dev") class Flux2DevRunner(Flux2BaseRunner): + def _rank0_text_encoder_broadcast_enabled(self): + return self.config.get("text_encoder_mode") == "rank0_broadcast" and dist.is_initialized() and dist.get_world_size() > 1 + def load_transformer(self): model_kwargs = { "model_path": os.path.join(self.config["model_path"], "transformer"), @@ -459,10 +510,33 @@ def init_scheduler(self): @ProfilingContext4DebugL1("Run Text Encoder") def run_text_encoder(self, text, image_list=None, neg_prompt=None): + if self._rank0_text_encoder_broadcast_enabled(): + return self._run_text_encoder_rank0_broadcast(text) + return self._run_text_encoder_local(text) + + def _encode_prompt(self, text): prompt_embeds_list, _ = self.text_encoders[0].infer([text]) - prompt_embeds = prompt_embeds_list[0].unsqueeze(0) + return prompt_embeds_list[0].unsqueeze(0) + + def _run_text_encoder_local(self, text): + prompt_embeds = self._encode_prompt(text) text_ids = self._prepare_text_ids(prompt_embeds).to(AI_DEVICE) - text_encoder_output = {"prompt_embeds": prompt_embeds, "text_ids": text_ids} + return {"prompt_embeds": prompt_embeds, "text_ids": text_ids} - return text_encoder_output + def _run_text_encoder_rank0_broadcast(self, text): + rank = dist.get_rank() + if rank == 0: + logger.info("[Flux2] Running Text Encoder on global rank 0 and broadcasting embeddings") + prompt_embeds = self._encode_prompt(text).contiguous() + prompt_shape = torch.tensor(prompt_embeds.shape, dtype=torch.long, device=prompt_embeds.device) + else: + prompt_shape = torch.empty(3, dtype=torch.long, device=AI_DEVICE) + + dist.broadcast(prompt_shape, src=0, group=dist.group.WORLD) + if rank != 0: + prompt_embeds = torch.empty(tuple(prompt_shape.tolist()), dtype=GET_DTYPE(), device=AI_DEVICE) + dist.broadcast(prompt_embeds, src=0, group=dist.group.WORLD) + + text_ids = self._prepare_text_ids(prompt_embeds).to(AI_DEVICE) + return {"prompt_embeds": prompt_embeds, "text_ids": text_ids} diff --git a/lightx2v/models/video_encoders/hf/flux2/vae.py b/lightx2v/models/video_encoders/hf/flux2/vae.py index f5a597332..401359e2e 100755 --- a/lightx2v/models/video_encoders/hf/flux2/vae.py +++ b/lightx2v/models/video_encoders/hf/flux2/vae.py @@ -53,12 +53,13 @@ def encode_vae_image(self, image): return encoded @torch.no_grad() - def decode(self, latents, input_info=None): + def decode(self, latents, input_info=None, output_type=None): if self.cpu_offload: self.vae.to(AI_DEVICE) image = self.vae.decode(latents.to(AI_DEVICE, dtype=GET_DTYPE()))[0] - output_type = "pt" if input_info is not None and getattr(input_info, "return_result_tensor", False) else "pil" + if output_type is None: + output_type = "pt" if input_info is not None and getattr(input_info, "return_result_tensor", False) else "pil" image = self.image_processor.postprocess(image, output_type=output_type) if self.cpu_offload: From d2aebd537c607d383d948be9937f4db24bacd9bb Mon Sep 17 00:00:00 2001 From: Super User Date: Sat, 8 Aug 2026 23:27:22 +0800 Subject: [PATCH 2/3] refactor(flux2): simplify distributed runner paths --- lightx2v/models/runners/flux2/flux2_runner.py | 48 ++++++++----------- 1 file changed, 20 insertions(+), 28 deletions(-) diff --git a/lightx2v/models/runners/flux2/flux2_runner.py b/lightx2v/models/runners/flux2/flux2_runner.py index 7205df684..3b248d59e 100644 --- a/lightx2v/models/runners/flux2/flux2_runner.py +++ b/lightx2v/models/runners/flux2/flux2_runner.py @@ -50,7 +50,10 @@ def _get_scheduler_class(self): @ProfilingContext4DebugL2("Load models") def load_model(self): - self.text_encoders = self._load_text_encoder_on_this_rank() + if self._rank0_text_encoder_broadcast_enabled() and dist.get_rank() != 0: + self.text_encoders = None + else: + self.text_encoders = self.load_text_encoder() self.vae = self.load_vae() if self._rank0_text_encoder_broadcast_enabled(): torch_device_module.synchronize() @@ -63,11 +66,6 @@ def load_model(self): def _rank0_text_encoder_broadcast_enabled(self): return False - def _load_text_encoder_on_this_rank(self): - if self._rank0_text_encoder_broadcast_enabled() and dist.get_rank() != 0: - return None - return self.load_text_encoder() - def load_vae(self): return Flux2VAE(self.config) @@ -91,9 +89,9 @@ def init_modules(self): def _run_input_encoder_local_t2i(self): prompt = self.input_info.prompt if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): - self.text_encoders = self._load_text_encoder_on_this_rank() + self.text_encoders = self.load_text_encoder() text_encoder_output = self.run_text_encoder(prompt, neg_prompt=self.input_info.negative_prompt) - if (self.config.get("lazy_load", False) or self.config.get("unload_modules", False)) and self.text_encoders is not None: + if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): del self.text_encoders[0] torch_device_module.empty_cache() gc.collect() @@ -106,9 +104,9 @@ def _run_input_encoder_local_t2i(self): def _run_input_encoder_local_i2i(self): prompt = self.input_info.prompt if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): - self.text_encoders = self._load_text_encoder_on_this_rank() + self.text_encoders = self.load_text_encoder() text_encoder_output = self.run_text_encoder(prompt, neg_prompt=self.input_info.negative_prompt) - if (self.config.get("lazy_load", False) or self.config.get("unload_modules", False)) and self.text_encoders is not None: + if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): del self.text_encoders[0] image_path = self.input_info.image_path @@ -354,7 +352,8 @@ def run_vae_decoder(self, latents): if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): self.vae = self.load_vae() - decode_rank = self._distributed_vae_decode_rank() + # vae_decode_rank is used only by distributed configs. + decode_rank = self.config.get("vae_decode_rank") if decode_rank is None: images = self._decode_latents_with_vae(latents) else: @@ -367,11 +366,6 @@ def run_vae_decoder(self, latents): return images - def _distributed_vae_decode_rank(self): - if not dist.is_initialized() or dist.get_world_size() == 1: - return None - return self.config.get("vae_decode_rank") - def _decode_vae_on_rank(self, latents, decode_rank): torch_device_module.synchronize() torch_device_module.empty_cache() @@ -384,10 +378,10 @@ def _decode_vae_on_rank(self, latents, decode_rank): else: image_shape = torch.empty(4, dtype=torch.long, device=latents.device) - dist.broadcast(image_shape, src=decode_rank, group=dist.group.WORLD) + dist.broadcast(image_shape, src=decode_rank) if rank != decode_rank: raw_images = torch.empty(tuple(image_shape.tolist()), dtype=GET_DTYPE(), device=latents.device) - dist.broadcast(raw_images, src=decode_rank, group=dist.group.WORLD) + dist.broadcast(raw_images, src=decode_rank) if self.input_info.return_result_tensor: return self.vae.image_processor.postprocess(raw_images, output_type="pt") @@ -485,7 +479,8 @@ def run_text_encoder(self, text, image_list=None, neg_prompt=None): @RUNNER_REGISTER("flux2_dev") class Flux2DevRunner(Flux2BaseRunner): def _rank0_text_encoder_broadcast_enabled(self): - return self.config.get("text_encoder_mode") == "rank0_broadcast" and dist.is_initialized() and dist.get_world_size() > 1 + # rank0_broadcast is used only by initialized multi-rank runs. + return self.config.get("text_encoder_mode") == "rank0_broadcast" def load_transformer(self): model_kwargs = { @@ -512,18 +507,15 @@ def init_scheduler(self): def run_text_encoder(self, text, image_list=None, neg_prompt=None): if self._rank0_text_encoder_broadcast_enabled(): return self._run_text_encoder_rank0_broadcast(text) - return self._run_text_encoder_local(text) - - def _encode_prompt(self, text): - prompt_embeds_list, _ = self.text_encoders[0].infer([text]) - return prompt_embeds_list[0].unsqueeze(0) - - def _run_text_encoder_local(self, text): prompt_embeds = self._encode_prompt(text) text_ids = self._prepare_text_ids(prompt_embeds).to(AI_DEVICE) return {"prompt_embeds": prompt_embeds, "text_ids": text_ids} + def _encode_prompt(self, text): + prompt_embeds_list, _ = self.text_encoders[0].infer([text]) + return prompt_embeds_list[0].unsqueeze(0) + def _run_text_encoder_rank0_broadcast(self, text): rank = dist.get_rank() if rank == 0: @@ -533,10 +525,10 @@ def _run_text_encoder_rank0_broadcast(self, text): else: prompt_shape = torch.empty(3, dtype=torch.long, device=AI_DEVICE) - dist.broadcast(prompt_shape, src=0, group=dist.group.WORLD) + dist.broadcast(prompt_shape, src=0) if rank != 0: prompt_embeds = torch.empty(tuple(prompt_shape.tolist()), dtype=GET_DTYPE(), device=AI_DEVICE) - dist.broadcast(prompt_embeds, src=0, group=dist.group.WORLD) + dist.broadcast(prompt_embeds, src=0) text_ids = self._prepare_text_ids(prompt_embeds).to(AI_DEVICE) return {"prompt_embeds": prompt_embeds, "text_ids": text_ids} From 229a9c4999d224ea2b86463cae33069bf6bbcf2b Mon Sep 17 00:00:00 2001 From: Super User Date: Sat, 8 Aug 2026 23:41:25 +0800 Subject: [PATCH 3/3] refactor(flux2): simplify text encoder mode checks --- lightx2v/models/runners/flux2/flux2_runner.py | 13 +++---------- 1 file changed, 3 insertions(+), 10 deletions(-) diff --git a/lightx2v/models/runners/flux2/flux2_runner.py b/lightx2v/models/runners/flux2/flux2_runner.py index 3b248d59e..c3d664ca7 100644 --- a/lightx2v/models/runners/flux2/flux2_runner.py +++ b/lightx2v/models/runners/flux2/flux2_runner.py @@ -50,12 +50,12 @@ def _get_scheduler_class(self): @ProfilingContext4DebugL2("Load models") def load_model(self): - if self._rank0_text_encoder_broadcast_enabled() and dist.get_rank() != 0: + if self.config.get("text_encoder_mode", False) == "rank0_broadcast" and dist.get_rank() != 0: self.text_encoders = None else: self.text_encoders = self.load_text_encoder() self.vae = self.load_vae() - if self._rank0_text_encoder_broadcast_enabled(): + if self.config.get("text_encoder_mode", False) == "rank0_broadcast": torch_device_module.synchronize() if AI_DEVICE == "cuda": dist.barrier(device_ids=[torch.cuda.current_device()]) @@ -63,9 +63,6 @@ def load_model(self): dist.barrier() self.model = self.load_transformer() - def _rank0_text_encoder_broadcast_enabled(self): - return False - def load_vae(self): return Flux2VAE(self.config) @@ -478,10 +475,6 @@ def run_text_encoder(self, text, image_list=None, neg_prompt=None): @RUNNER_REGISTER("flux2_dev") class Flux2DevRunner(Flux2BaseRunner): - def _rank0_text_encoder_broadcast_enabled(self): - # rank0_broadcast is used only by initialized multi-rank runs. - return self.config.get("text_encoder_mode") == "rank0_broadcast" - def load_transformer(self): model_kwargs = { "model_path": os.path.join(self.config["model_path"], "transformer"), @@ -505,7 +498,7 @@ def init_scheduler(self): @ProfilingContext4DebugL1("Run Text Encoder") def run_text_encoder(self, text, image_list=None, neg_prompt=None): - if self._rank0_text_encoder_broadcast_enabled(): + if self.config.get("text_encoder_mode", False) == "rank0_broadcast": return self._run_text_encoder_rank0_broadcast(text) prompt_embeds = self._encode_prompt(text) text_ids = self._prepare_text_ids(prompt_embeds).to(AI_DEVICE)