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: 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