[diffusion] fix: respect configured precision in Qwen layered path (#21980)

Co-authored-by: jxp <jingxin.pan123@gmail.com>
This commit is contained in:
jy-song-hub
2026-05-20 10:38:03 +08:00
committed by GitHub
co-authored by jxp
parent 3a9d9d5832
commit 549ae16c6f
2 changed files with 25 additions and 6 deletions
@@ -13,6 +13,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.q
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
# TODO(will): move PRECISION_TO_TYPE to better place
@@ -122,6 +123,10 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
model_path=self.model_path,
vae_dtype=PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision],
text_encoder_dtype=PRECISION_TO_TYPE[
server_args.pipeline_config.text_encoder_precisions[0]
],
)
)
@@ -123,15 +123,29 @@ def retrieve_timesteps(
class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
def __init__(
self, vae, tokenizer, processor, transformer, scheduler, model_path
self,
vae,
tokenizer,
processor,
transformer,
scheduler,
model_path,
vae_dtype: torch.dtype,
text_encoder_dtype: torch.dtype,
) -> None:
super().__init__()
self.vae = vae.to(torch.bfloat16)
self.vae = vae.to(dtype=vae_dtype)
self.vae_dtype = vae_dtype
self.text_encoder_dtype = text_encoder_dtype
from transformers import Qwen2_5_VLForConditionalGeneration
self.text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained(
model_path, subfolder="text_encoder"
).to(torch.bfloat16)
self.text_encoder = (
Qwen2_5_VLForConditionalGeneration.from_pretrained(
model_path, subfolder="text_encoder"
)
.to(get_local_torch_device())
.to(dtype=self.text_encoder_dtype)
)
self.tokenizer = tokenizer
self.processor = processor
self.transformer = transformer
@@ -464,7 +478,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
image, calculated_height, calculated_width
)
image = image.unsqueeze(2)
image = image.to(dtype=torch.bfloat16)
image = image.to(dtype=self.vae_dtype)
prompt = batch.prompt
with self.use_declared_component(