From 681ceba58162d92d0aaadfb9c03996756665cc71 Mon Sep 17 00:00:00 2001 From: rockdu Date: Sat, 8 Aug 2026 01:56:29 -0700 Subject: [PATCH 1/2] feat(rollout): forward --sglang-text-encoder-precisions to the engine --- .../backends/sglang_diffusion_utils/sglang_diffusion_engine.py | 2 +- miles/utils/arguments.py | 1 + 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/miles/backends/sglang_diffusion_utils/sglang_diffusion_engine.py b/miles/backends/sglang_diffusion_utils/sglang_diffusion_engine.py index 3de4dd1a..1f22bdc6 100644 --- a/miles/backends/sglang_diffusion_utils/sglang_diffusion_engine.py +++ b/miles/backends/sglang_diffusion_utils/sglang_diffusion_engine.py @@ -340,7 +340,7 @@ def _compute_server_args(args, host, port, nccl_port): # dit_precision / vae_precision are PipelineConfig fields, not ServerArgs, so forward them explicitly (only when changed from the class default, to avoid clobbering a subclass override). from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig - for field_name in ("dit_precision", "vae_precision"): + for field_name in ("dit_precision", "vae_precision", "text_encoder_precisions"): val = getattr(args, f"sglang_{field_name}", None) if val is not None and val != getattr(PipelineConfig, field_name, None): kwargs[field_name] = val diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 5a6821de..2a1698c8 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -1420,6 +1420,7 @@ def add_sglang_tp_size(): ) parser.set_defaults(sglang_tensor_parallel_size=add_sglang_tp_size()) + parser.set_defaults(sglang_text_encoder_precisions=None) return parser return add_miles_arguments From 3e458bb7f87c2c7352ddce9cc16063469d34544d Mon Sep 17 00:00:00 2001 From: rockdu Date: Sun, 9 Aug 2026 01:38:43 -0700 Subject: [PATCH 2/2] feat(sd3): mirror the rollout boundary input dtypes --- miles/backends/fsdp_utils/configs/sd3.py | 1 + 1 file changed, 1 insertion(+) diff --git a/miles/backends/fsdp_utils/configs/sd3.py b/miles/backends/fsdp_utils/configs/sd3.py index 332e9fa1..202ea5fa 100644 --- a/miles/backends/fsdp_utils/configs/sd3.py +++ b/miles/backends/fsdp_utils/configs/sd3.py @@ -26,6 +26,7 @@ class SD3TrainPipelineConfig(TrainPipelineConfig): "attn.to_add_out", ] needs_timestep_scaling = False + input_dtype_policy = {"latents": "default", "cond": None, "timestep": "fp32"} def prepare_cond_kwargs(self, cond: CondKwargs | None, device: torch.device) -> dict: if cond is None: