[diffusion] fix: add precision consistency layer (#27088)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -17,7 +17,7 @@ from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
|||||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
||||||
get_diffusers_component_config,
|
get_diffusers_component_config,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||||
|
|
||||||
|
|
||||||
class AdapterLoader(ComponentLoader):
|
class AdapterLoader(ComponentLoader):
|
||||||
@@ -51,7 +51,9 @@ class AdapterLoader(ComponentLoader):
|
|||||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||||
|
|
||||||
target_device = get_local_torch_device()
|
target_device = get_local_torch_device()
|
||||||
default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
|
default_dtype = resolve_precision(
|
||||||
|
server_args, "connectors", precision_attr="dit_precision"
|
||||||
|
)
|
||||||
|
|
||||||
with set_default_torch_dtype(default_dtype), skip_init_modules():
|
with set_default_torch_dtype(default_dtype), skip_init_modules():
|
||||||
connector_cfg = LTX2ConnectorConfig()
|
connector_cfg = LTX2ConnectorConfig()
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
|||||||
get_diffusers_component_config,
|
get_diffusers_component_config,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -62,7 +62,9 @@ class BridgeLoader(ComponentLoader):
|
|||||||
if not safetensors_list:
|
if not safetensors_list:
|
||||||
raise ValueError(f"No safetensors files found in {component_model_path}")
|
raise ValueError(f"No safetensors files found in {component_model_path}")
|
||||||
|
|
||||||
default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
|
default_dtype = resolve_precision(
|
||||||
|
server_args, component_name, precision_attr="dit_precision"
|
||||||
|
)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Loading %s from %s safetensors files, default_dtype: %s",
|
"Loading %s from %s safetensors files, default_dtype: %s",
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
|||||||
prepare_diffusers_component_path_for_loading,
|
prepare_diffusers_component_path_for_loading,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_component_precision
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -232,7 +233,10 @@ class ComponentLoader(ABC):
|
|||||||
attn_backend, component_name=component_attn_name
|
attn_backend, component_name=component_attn_name
|
||||||
):
|
):
|
||||||
component = self.load_native(
|
component = self.load_native(
|
||||||
component_model_path, server_args, transformers_or_diffusers
|
component_model_path,
|
||||||
|
server_args,
|
||||||
|
transformers_or_diffusers,
|
||||||
|
component_name,
|
||||||
)
|
)
|
||||||
should_offload = self.should_offload(server_args)
|
should_offload = self.should_offload(server_args)
|
||||||
target_device = self.target_device(should_offload)
|
target_device = self.target_device(should_offload)
|
||||||
@@ -268,10 +272,20 @@ class ComponentLoader(ABC):
|
|||||||
component_model_path: str,
|
component_model_path: str,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
transformers_or_diffusers: str,
|
transformers_or_diffusers: str,
|
||||||
|
component_name: str | None = None,
|
||||||
) -> AutoModel:
|
) -> AutoModel:
|
||||||
"""
|
"""
|
||||||
Load the component using the native library (transformers/diffusers).
|
Load the component using the native library (transformers/diffusers).
|
||||||
"""
|
"""
|
||||||
|
precision = (
|
||||||
|
resolve_component_precision(server_args, component_name)
|
||||||
|
if component_name is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
load_kwargs = {}
|
||||||
|
if precision is not None:
|
||||||
|
load_kwargs["torch_dtype"] = precision
|
||||||
|
|
||||||
if transformers_or_diffusers == "transformers":
|
if transformers_or_diffusers == "transformers":
|
||||||
from transformers import AutoModel
|
from transformers import AutoModel
|
||||||
|
|
||||||
@@ -285,6 +299,7 @@ class ComponentLoader(ABC):
|
|||||||
config=config,
|
config=config,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
|
**load_kwargs,
|
||||||
)
|
)
|
||||||
elif transformers_or_diffusers == "diffusers":
|
elif transformers_or_diffusers == "diffusers":
|
||||||
from diffusers import AutoModel
|
from diffusers import AutoModel
|
||||||
@@ -296,6 +311,7 @@ class ComponentLoader(ABC):
|
|||||||
component_model_path,
|
component_model_path,
|
||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
|
**load_kwargs,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported library: {transformers_or_diffusers}")
|
raise ValueError(f"Unsupported library: {transformers_or_diffusers}")
|
||||||
|
|||||||
+14
-3
@@ -38,6 +38,7 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
|||||||
get_diffusers_component_config,
|
get_diffusers_component_config,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import precision_to_dtype
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
@@ -98,16 +99,26 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
component_model_path: str,
|
component_model_path: str,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
transformers_or_diffusers: str,
|
transformers_or_diffusers: str,
|
||||||
|
component_name: str | None = None,
|
||||||
):
|
):
|
||||||
if transformers_or_diffusers != "transformers":
|
if transformers_or_diffusers != "transformers":
|
||||||
return super().load_native(
|
return super().load_native(
|
||||||
component_model_path, server_args, transformers_or_diffusers
|
component_model_path,
|
||||||
|
server_args,
|
||||||
|
transformers_or_diffusers,
|
||||||
|
component_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
encoder_idx = (
|
encoder_idx = (
|
||||||
1 if component_model_path.rstrip("/").endswith("text_encoder_2") else 0
|
self._extract_encoder_index(component_name or "text_encoder_2")
|
||||||
|
if component_name
|
||||||
|
else 1 if component_model_path.rstrip("/").endswith("text_encoder_2") else 0
|
||||||
)
|
)
|
||||||
encoder_dtype = server_args.pipeline_config.text_encoder_precisions[encoder_idx]
|
encoder_dtype = server_args.pipeline_config.text_encoder_precisions[encoder_idx]
|
||||||
|
dtype = precision_to_dtype(
|
||||||
|
encoder_dtype,
|
||||||
|
f"text_encoder_precisions[{encoder_idx}]",
|
||||||
|
)
|
||||||
transformers_model_class = self._resolve_transformers_text_encoder_class(
|
transformers_model_class = self._resolve_transformers_text_encoder_class(
|
||||||
component_model_path, server_args
|
component_model_path, server_args
|
||||||
)
|
)
|
||||||
@@ -115,7 +126,7 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
component_model_path,
|
component_model_path,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
torch_dtype=PRECISION_TO_TYPE[encoder_dtype],
|
torch_dtype=dtype,
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
|||||||
get_diffusers_component_config,
|
get_diffusers_component_config,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_component_precision
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -127,6 +128,12 @@ class VAELoader(ComponentLoader):
|
|||||||
)
|
)
|
||||||
vae_config = getattr(server_args.pipeline_config, pipeline_vae_config_attr)
|
vae_config = getattr(server_args.pipeline_config, pipeline_vae_config_attr)
|
||||||
vae_precision = getattr(server_args.pipeline_config, pipeline_vae_precision)
|
vae_precision = getattr(server_args.pipeline_config, pipeline_vae_precision)
|
||||||
|
resolved_vae_dtype = resolve_component_precision(server_args, component_name)
|
||||||
|
vae_dtype = (
|
||||||
|
resolved_vae_dtype
|
||||||
|
if resolved_vae_dtype is not None
|
||||||
|
else PRECISION_TO_TYPE[vae_precision]
|
||||||
|
)
|
||||||
vae_config.update_model_arch(config)
|
vae_config.update_model_arch(config)
|
||||||
if hasattr(vae_config, "post_init"):
|
if hasattr(vae_config, "post_init"):
|
||||||
# NOTE: some post init logics are only available after updated with config
|
# NOTE: some post init logics are only available after updated with config
|
||||||
@@ -145,7 +152,6 @@ class VAELoader(ComponentLoader):
|
|||||||
custom_module = importlib.util.module_from_spec(spec)
|
custom_module = importlib.util.module_from_spec(spec)
|
||||||
spec.loader.exec_module(custom_module)
|
spec.loader.exec_module(custom_module)
|
||||||
vae_cls = getattr(custom_module, cls_name)
|
vae_cls = getattr(custom_module, cls_name)
|
||||||
vae_dtype = PRECISION_TO_TYPE[vae_precision]
|
|
||||||
with set_default_torch_dtype(vae_dtype):
|
with set_default_torch_dtype(vae_dtype):
|
||||||
vae = vae_cls.from_pretrained(
|
vae = vae_cls.from_pretrained(
|
||||||
component_model_path,
|
component_model_path,
|
||||||
@@ -164,7 +170,7 @@ class VAELoader(ComponentLoader):
|
|||||||
|
|
||||||
# Load from ModelRegistry (standard VAE classes)
|
# Load from ModelRegistry (standard VAE classes)
|
||||||
with (
|
with (
|
||||||
set_default_torch_dtype(PRECISION_TO_TYPE[vae_precision]),
|
set_default_torch_dtype(vae_dtype),
|
||||||
skip_init_modules(),
|
skip_init_modules(),
|
||||||
):
|
):
|
||||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
|||||||
get_diffusers_component_config,
|
get_diffusers_component_config,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_component_precision
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -47,11 +48,12 @@ class VocoderLoader(ComponentLoader):
|
|||||||
vocoder_config = LTXVocoderConfig()
|
vocoder_config = LTXVocoderConfig()
|
||||||
vocoder_config.update_model_arch(config)
|
vocoder_config.update_model_arch(config)
|
||||||
|
|
||||||
try:
|
resolved_vocoder_dtype = resolve_component_precision(server_args, "vocoder")
|
||||||
vocoder_precision = server_args.pipeline_config.audio_vae_precision
|
vocoder_dtype = (
|
||||||
except AttributeError:
|
resolved_vocoder_dtype
|
||||||
vocoder_precision = "fp32"
|
if resolved_vocoder_dtype is not None
|
||||||
vocoder_dtype = PRECISION_TO_TYPE[vocoder_precision]
|
else PRECISION_TO_TYPE["fp32"]
|
||||||
|
)
|
||||||
|
|
||||||
should_offload = self.should_offload(server_args)
|
should_offload = self.should_offload(server_args)
|
||||||
target_device = self.target_device(should_offload)
|
target_device = self.target_device(should_offload)
|
||||||
|
|||||||
@@ -24,13 +24,13 @@ from sglang.multimodal_gen.runtime.loader.utils import _list_safetensors_files
|
|||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||||
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
||||||
build_nvfp4_config_from_safetensors_list,
|
build_nvfp4_config_from_safetensors_list,
|
||||||
get_metadata_from_safetensors_file,
|
get_metadata_from_safetensors_file,
|
||||||
get_quant_config,
|
get_quant_config,
|
||||||
get_quant_config_from_safetensors_metadata,
|
get_quant_config_from_safetensors_metadata,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
|
||||||
from sglang.srt.layers.quantization import QuantizationConfig
|
from sglang.srt.layers.quantization import QuantizationConfig
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -617,4 +617,4 @@ def _resolve_target_param_dtype(
|
|||||||
) -> Optional[torch.dtype]:
|
) -> Optional[torch.dtype]:
|
||||||
if quant_config is not None or nunchaku_config is not None:
|
if quant_config is not None or nunchaku_config is not None:
|
||||||
return None
|
return None
|
||||||
return PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
|
return resolve_precision(server_args, "dit", precision_attr="dit_precision")
|
||||||
|
|||||||
@@ -956,8 +956,7 @@ class AutoencoderKLQwenImage(ParallelTiledVAE):
|
|||||||
else 0,
|
else 0,
|
||||||
}
|
}
|
||||||
cuda_device = get_local_torch_device()
|
cuda_device = get_local_torch_device()
|
||||||
# FIXME: hardcode
|
dtype = torch.get_default_dtype()
|
||||||
dtype = torch.bfloat16
|
|
||||||
latent_channels = config.arch_config.z_dim
|
latent_channels = config.arch_config.z_dim
|
||||||
|
|
||||||
self.shift_factor = (
|
self.shift_factor = (
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -613,7 +613,9 @@ class ComfyUIFluxPipeline(LoRAPipeline, ComposedPipelineBase):
|
|||||||
|
|
||||||
safetensors_list = [self.model_path]
|
safetensors_list = [self.model_path]
|
||||||
logger.info("Loading weights from: %s", safetensors_list)
|
logger.info("Loading weights from: %s", safetensors_list)
|
||||||
default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
|
default_dtype = resolve_precision(
|
||||||
|
server_args, "dit", precision_attr="dit_precision"
|
||||||
|
)
|
||||||
server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
|
server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
|
||||||
hf_config = {}
|
hf_config = {}
|
||||||
|
|
||||||
|
|||||||
@@ -35,7 +35,8 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE, set_mixed_precision_policy
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||||
|
from sglang.multimodal_gen.utils import set_mixed_precision_policy
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -185,7 +186,9 @@ class ComfyUIQwenImagePipelineBase(LoRAPipeline, ComposedPipelineBase):
|
|||||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||||
logger.info("Resolved transformer class: %s", cls_name)
|
logger.info("Resolved transformer class: %s", cls_name)
|
||||||
|
|
||||||
default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
|
default_dtype = resolve_precision(
|
||||||
|
server_args, "dit", precision_attr="dit_precision"
|
||||||
|
)
|
||||||
server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
|
server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
|
||||||
assert server_args.hsdp_shard_dim is not None, "hsdp_shard_dim must be set"
|
assert server_args.hsdp_shard_dim is not None, "hsdp_shard_dim must be set"
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -214,6 +217,8 @@ class ComfyUIQwenImagePipelineBase(LoRAPipeline, ComposedPipelineBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
# precision-constraint: FSDP mixed precision currently uses bf16
|
||||||
|
# parameters and fp32 reduction regardless of model load dtype.
|
||||||
mp_policy = MixedPrecisionPolicy(
|
mp_policy = MixedPrecisionPolicy(
|
||||||
torch.bfloat16, torch.float32, None, cast_forward_inputs=False
|
torch.bfloat16, torch.float32, None, cast_forward_inputs=False
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -38,7 +38,8 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE, set_mixed_precision_policy
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||||
|
from sglang.multimodal_gen.utils import set_mixed_precision_policy
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -267,7 +268,9 @@ class ComfyUIZImagePipeline(LoRAPipeline, ComposedPipelineBase):
|
|||||||
safetensors_list = [self.model_path]
|
safetensors_list = [self.model_path]
|
||||||
logger.info("Loading weights from: %s", safetensors_list)
|
logger.info("Loading weights from: %s", safetensors_list)
|
||||||
|
|
||||||
default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
|
default_dtype = resolve_precision(
|
||||||
|
server_args, "dit", precision_attr="dit_precision"
|
||||||
|
)
|
||||||
server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
|
server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
|
||||||
hf_config = {}
|
hf_config = {}
|
||||||
|
|
||||||
@@ -289,6 +292,8 @@ class ComfyUIZImagePipeline(LoRAPipeline, ComposedPipelineBase):
|
|||||||
# Create model first (same as maybe_load_fsdp_model)
|
# Create model first (same as maybe_load_fsdp_model)
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
|
|
||||||
|
# precision-constraint: FSDP mixed precision currently uses bf16
|
||||||
|
# parameters and fp32 reduction regardless of model load dtype.
|
||||||
mp_policy = MixedPrecisionPolicy(
|
mp_policy = MixedPrecisionPolicy(
|
||||||
torch.bfloat16, torch.float32, None, cast_forward_inputs=False
|
torch.bfloat16, torch.float32, None, cast_forward_inputs=False
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ from sglang.multimodal_gen.runtime.platforms import (
|
|||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -667,21 +668,15 @@ class DiffusersPipeline(ComposedPipelineBase):
|
|||||||
return pipe
|
return pipe
|
||||||
|
|
||||||
def _get_dtype(self, server_args: ServerArgs) -> torch.dtype:
|
def _get_dtype(self, server_args: ServerArgs) -> torch.dtype:
|
||||||
dtype = (
|
"""
|
||||||
torch.bfloat16
|
Determine the dtype to use for model loading.
|
||||||
if torch.get_device_module().is_bf16_supported()
|
"""
|
||||||
else torch.float16
|
if hasattr(server_args, "pipeline_config") and server_args.pipeline_config:
|
||||||
)
|
return resolve_precision(server_args, "dit", precision_attr="dit_precision")
|
||||||
|
|
||||||
dit_precision = server_args.pipeline_config.dit_precision
|
# precision-constraint: legacy fallback for callers without pipeline_config;
|
||||||
if dit_precision == "fp16":
|
# prefer explicit dit_precision policy when available.
|
||||||
dtype = torch.float16
|
return torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
|
||||||
elif dit_precision == "bf16":
|
|
||||||
dtype = torch.bfloat16
|
|
||||||
elif dit_precision == "fp32":
|
|
||||||
dtype = torch.float32
|
|
||||||
|
|
||||||
return dtype
|
|
||||||
|
|
||||||
def _detect_pipeline_type(self) -> None:
|
def _detect_pipeline_type(self) -> None:
|
||||||
"""Detect if this is an image or video pipeline."""
|
"""Detect if this is an image or video pipeline."""
|
||||||
|
|||||||
@@ -35,6 +35,9 @@ from sglang.multimodal_gen.runtime.realtime.states import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||||
|
autocast_enabled as precision_autocast_enabled,
|
||||||
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -114,7 +117,9 @@ class CausalDMDDenoisingStage(DenoisingStage):
|
|||||||
target_dtype: torch.dtype,
|
target_dtype: torch.dtype,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
return (target_dtype != torch.float32) and not server_args.disable_autocast
|
# precision-constraint: Causal denoising kernels are validated on bf16;
|
||||||
|
# do not replace this with user precision policy without auditing kernel support.
|
||||||
|
return precision_autocast_enabled(target_dtype, server_args.disable_autocast)
|
||||||
|
|
||||||
def _prepare_frame_seq_length(self, h: int, w: int) -> int:
|
def _prepare_frame_seq_length(self, h: int, w: int) -> int:
|
||||||
patch_ratio = (
|
patch_ratio = (
|
||||||
|
|||||||
@@ -30,7 +30,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
|||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||||
|
autocast_enabled,
|
||||||
|
resolve_precision,
|
||||||
|
temporary_module_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -105,7 +109,9 @@ class DecodingStage(PipelineStage):
|
|||||||
def component_uses(
|
def component_uses(
|
||||||
self, server_args: ServerArgs, stage_name: str | None = None
|
self, server_args: ServerArgs, stage_name: str | None = None
|
||||||
) -> list[ComponentUse]:
|
) -> list[ComponentUse]:
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
|
server_args, self.component_name, precision_attr="vae_precision"
|
||||||
|
)
|
||||||
stage_name = self._component_stage_name(stage_name)
|
stage_name = self._component_stage_name(stage_name)
|
||||||
return [
|
return [
|
||||||
ComponentUse(
|
ComponentUse(
|
||||||
@@ -172,9 +178,11 @@ class DecodingStage(PipelineStage):
|
|||||||
normalized to [0, 1] range and moved to CPU as float32
|
normalized to [0, 1] range and moved to CPU as float32
|
||||||
"""
|
"""
|
||||||
latents = latents.to(get_local_torch_device())
|
latents = latents.to(get_local_torch_device())
|
||||||
vae_autocast_enabled = (
|
# Setup VAE precision from user policy.
|
||||||
vae_dtype != torch.float32
|
vae_dtype = resolve_precision(
|
||||||
) and not server_args.disable_autocast
|
server_args, self.component_name, precision_attr="vae_precision"
|
||||||
|
)
|
||||||
|
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
|
||||||
|
|
||||||
# scale and shift
|
# scale and shift
|
||||||
latents = self.scale_and_shift(latents, server_args)
|
latents = self.scale_and_shift(latents, server_args)
|
||||||
@@ -195,10 +203,14 @@ class DecodingStage(PipelineStage):
|
|||||||
self.vae.enable_tiling()
|
self.vae.enable_tiling()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
should_cast_vae = not vae_autocast_enabled
|
||||||
if not vae_autocast_enabled:
|
if not vae_autocast_enabled:
|
||||||
latents = latents.to(vae_dtype)
|
latents = latents.to(vae_dtype)
|
||||||
decode_output = self.vae.decode(latents)
|
with temporary_module_dtype(
|
||||||
image = _ensure_tensor_decode_output(decode_output)
|
self.vae, vae_dtype, enabled=should_cast_vae
|
||||||
|
) as vae:
|
||||||
|
decode_output = vae.decode(latents)
|
||||||
|
image = _ensure_tensor_decode_output(decode_output)
|
||||||
|
|
||||||
# De-normalize image to [0, 1] range
|
# De-normalize image to [0, 1] range
|
||||||
image = (image / 2 + 0.5).clamp(0, 1)
|
image = (image / 2 + 0.5).clamp(0, 1)
|
||||||
@@ -236,7 +248,9 @@ class DecodingStage(PipelineStage):
|
|||||||
# load vae if not already loaded (used for memory constrained devices)
|
# load vae if not already loaded (used for memory constrained devices)
|
||||||
self.load_model()
|
self.load_model()
|
||||||
|
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
|
server_args, self.component_name, precision_attr="vae_precision"
|
||||||
|
)
|
||||||
with self.use_declared_component(
|
with self.use_declared_component(
|
||||||
component_name=self.component_name,
|
component_name=self.component_name,
|
||||||
module=self.vae,
|
module=self.vae,
|
||||||
|
|||||||
@@ -101,8 +101,14 @@ from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
|||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import maybe_nvtx_range
|
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import maybe_nvtx_range
|
||||||
from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler
|
from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||||
|
autocast_enabled as precision_autocast_enabled,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||||
|
resolve_precision,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler
|
from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE, dict_to_3d_list
|
from sglang.multimodal_gen.utils import dict_to_3d_list
|
||||||
from sglang.srt.utils.common import get_compiler_backend
|
from sglang.srt.utils.common import get_compiler_backend
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -196,6 +202,8 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
self.pipeline = weakref.ref(pipeline) if pipeline else None
|
self.pipeline = weakref.ref(pipeline) if pipeline else None
|
||||||
|
|
||||||
selected_attention_backend = self._infer_transformer_attention_backend()
|
selected_attention_backend = self._infer_transformer_attention_backend()
|
||||||
|
# precision-constraint: attention backend metadata allocation currently assumes fp16;
|
||||||
|
# do not replace with user precision policy without auditing backend support.
|
||||||
self.attn_backend = get_attn_backend(
|
self.attn_backend = get_attn_backend(
|
||||||
head_size=attn_head_size,
|
head_size=attn_head_size,
|
||||||
dtype=torch.float16,
|
dtype=torch.float16,
|
||||||
@@ -244,7 +252,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
stage_name = self._component_stage_name(stage_name)
|
stage_name = self._component_stage_name(stage_name)
|
||||||
uses: list[ComponentUse] = []
|
uses: list[ComponentUse] = []
|
||||||
if self.vae is not None:
|
if self.vae is not None:
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
|
server_args, "vae", precision_attr="vae_precision"
|
||||||
|
)
|
||||||
uses.append(
|
uses.append(
|
||||||
ComponentUse(
|
ComponentUse(
|
||||||
stage_name=stage_name,
|
stage_name=stage_name,
|
||||||
@@ -656,10 +666,12 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Setup precision and autocast settings
|
# Setup precision and autocast settings
|
||||||
target_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
|
target_dtype = resolve_precision(
|
||||||
autocast_enabled = (
|
server_args, "dit", precision_attr="dit_precision"
|
||||||
target_dtype != torch.float32
|
)
|
||||||
) and not server_args.disable_autocast
|
autocast_enabled = precision_autocast_enabled(
|
||||||
|
target_dtype, server_args.disable_autocast
|
||||||
|
)
|
||||||
|
|
||||||
# Prepare image latents and embeddings for I2V generation
|
# Prepare image latents and embeddings for I2V generation
|
||||||
image_embeds = batch.image_embeds
|
image_embeds = batch.image_embeds
|
||||||
@@ -683,7 +695,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
|
|
||||||
# TI2V specific preparations - before SP sharding
|
# TI2V specific preparations - before SP sharding
|
||||||
if should_preprocess_for_wan_ti2v:
|
if should_preprocess_for_wan_ti2v:
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
|
server_args, "vae", precision_attr="vae_precision"
|
||||||
|
)
|
||||||
with self.use_declared_component(
|
with self.use_declared_component(
|
||||||
component_name="vae",
|
component_name="vae",
|
||||||
module=self.vae,
|
module=self.vae,
|
||||||
|
|||||||
@@ -23,7 +23,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
|||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||||
|
autocast_enabled,
|
||||||
|
resolve_precision,
|
||||||
|
temporary_module_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -43,7 +47,9 @@ class EncodingStage(PipelineStage):
|
|||||||
def component_uses(
|
def component_uses(
|
||||||
self, server_args: ServerArgs, stage_name: str | None = None
|
self, server_args: ServerArgs, stage_name: str | None = None
|
||||||
) -> list[ComponentUse]:
|
) -> list[ComponentUse]:
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
|
server_args, "vae", precision_attr="vae_precision"
|
||||||
|
)
|
||||||
stage_name = self._component_stage_name(stage_name)
|
stage_name = self._component_stage_name(stage_name)
|
||||||
return [
|
return [
|
||||||
ComponentUse(
|
ComponentUse(
|
||||||
@@ -83,11 +89,11 @@ class EncodingStage(PipelineStage):
|
|||||||
"""
|
"""
|
||||||
assert batch.latents is not None and isinstance(batch.latents, torch.Tensor)
|
assert batch.latents is not None and isinstance(batch.latents, torch.Tensor)
|
||||||
|
|
||||||
# Setup VAE precision
|
# Setup VAE precision from user policy.
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
vae_autocast_enabled = (
|
server_args, "vae", precision_attr="vae_precision"
|
||||||
vae_dtype != torch.float32
|
)
|
||||||
) and not server_args.disable_autocast
|
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
|
||||||
|
|
||||||
# Normalize input to [-1, 1] range (reverse of decoding normalization)
|
# Normalize input to [-1, 1] range (reverse of decoding normalization)
|
||||||
latents = (batch.latents * 2.0 - 1.0).clamp(-1, 1)
|
latents = (batch.latents * 2.0 - 1.0).clamp(-1, 1)
|
||||||
@@ -109,9 +115,13 @@ class EncodingStage(PipelineStage):
|
|||||||
self.vae.enable_tiling()
|
self.vae.enable_tiling()
|
||||||
# if server_args.vae_sp:
|
# if server_args.vae_sp:
|
||||||
# self.vae.enable_parallel()
|
# self.vae.enable_parallel()
|
||||||
|
should_cast_vae = not vae_autocast_enabled
|
||||||
if not vae_autocast_enabled:
|
if not vae_autocast_enabled:
|
||||||
latents = latents.to(vae_dtype)
|
latents = latents.to(vae_dtype)
|
||||||
latents = self.vae.encode(latents).mean
|
with temporary_module_dtype(
|
||||||
|
self.vae, vae_dtype, enabled=should_cast_vae
|
||||||
|
) as vae:
|
||||||
|
latents = vae.encode(latents).mean
|
||||||
|
|
||||||
# Update batch with encoded latents
|
# Update batch with encoded latents
|
||||||
batch.latents = latents
|
batch.latents = latents
|
||||||
|
|||||||
@@ -44,7 +44,12 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
|||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||||
|
align_tensor_to_module_dtype,
|
||||||
|
autocast_enabled,
|
||||||
|
resolve_precision,
|
||||||
|
temporary_module_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -250,6 +255,14 @@ class ImageEncodingStage(PipelineStage):
|
|||||||
) as image_encoder:
|
) as image_encoder:
|
||||||
assert image_encoder is not None
|
assert image_encoder is not None
|
||||||
self.image_encoder = image_encoder
|
self.image_encoder = image_encoder
|
||||||
|
if hasattr(image_inputs, "pixel_values") and isinstance(
|
||||||
|
image_inputs.pixel_values, torch.Tensor
|
||||||
|
):
|
||||||
|
image_inputs["pixel_values"] = align_tensor_to_module_dtype(
|
||||||
|
image_inputs.pixel_values,
|
||||||
|
self.image_encoder,
|
||||||
|
device=cuda_device,
|
||||||
|
)
|
||||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||||
outputs = self.image_encoder(
|
outputs = self.image_encoder(
|
||||||
**image_inputs,
|
**image_inputs,
|
||||||
@@ -284,6 +297,24 @@ class ImageEncodingStage(PipelineStage):
|
|||||||
) as text_encoder:
|
) as text_encoder:
|
||||||
assert text_encoder is not None
|
assert text_encoder is not None
|
||||||
self.text_encoder = text_encoder
|
self.text_encoder = text_encoder
|
||||||
|
if hasattr(image_inputs, "pixel_values") and isinstance(
|
||||||
|
image_inputs.pixel_values, torch.Tensor
|
||||||
|
):
|
||||||
|
image_inputs["pixel_values"] = align_tensor_to_module_dtype(
|
||||||
|
image_inputs.pixel_values,
|
||||||
|
self.text_encoder,
|
||||||
|
device=cuda_device,
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
batch.do_classifier_free_guidance
|
||||||
|
and hasattr(neg_image_inputs, "pixel_values")
|
||||||
|
and isinstance(neg_image_inputs.pixel_values, torch.Tensor)
|
||||||
|
):
|
||||||
|
neg_image_inputs["pixel_values"] = align_tensor_to_module_dtype(
|
||||||
|
neg_image_inputs.pixel_values,
|
||||||
|
self.text_encoder,
|
||||||
|
device=cuda_device,
|
||||||
|
)
|
||||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||||
outputs = self.text_encoder(
|
outputs = self.text_encoder(
|
||||||
input_ids=image_inputs.input_ids,
|
input_ids=image_inputs.input_ids,
|
||||||
@@ -573,10 +604,10 @@ class LTX2ImageEncodingStage(PipelineStage):
|
|||||||
generator: torch.Generator | None,
|
generator: torch.Generator | None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""VAE encode → sample → per-channel normalize (LTX-2 convention)."""
|
"""VAE encode → sample → per-channel normalize (LTX-2 convention)."""
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
vae_autocast_enabled = (
|
server_args, "vae", precision_attr="vae_precision"
|
||||||
vae_dtype != torch.float32
|
)
|
||||||
) and not server_args.disable_autocast
|
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
|
||||||
|
|
||||||
with torch.autocast(
|
with torch.autocast(
|
||||||
device_type=current_platform.device_type,
|
device_type=current_platform.device_type,
|
||||||
@@ -588,7 +619,13 @@ class LTX2ImageEncodingStage(PipelineStage):
|
|||||||
self.vae.enable_tiling()
|
self.vae.enable_tiling()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
latent_dist = self.vae.encode(video_condition)
|
should_cast_vae = not vae_autocast_enabled
|
||||||
|
if not vae_autocast_enabled:
|
||||||
|
video_condition = video_condition.to(vae_dtype)
|
||||||
|
with temporary_module_dtype(
|
||||||
|
self.vae, vae_dtype, enabled=should_cast_vae
|
||||||
|
) as vae:
|
||||||
|
latent_dist = vae.encode(video_condition)
|
||||||
if isinstance(latent_dist, AutoencoderKLOutput):
|
if isinstance(latent_dist, AutoencoderKLOutput):
|
||||||
latent_dist = latent_dist.latent_dist
|
latent_dist = latent_dist.latent_dist
|
||||||
|
|
||||||
@@ -610,10 +647,10 @@ class LTX2ImageEncodingStage(PipelineStage):
|
|||||||
self, video_condition: torch.Tensor, server_args: ServerArgs
|
self, video_condition: torch.Tensor, server_args: ServerArgs
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""LTX-2.3 condition-image encoder path (bypasses VAE)."""
|
"""LTX-2.3 condition-image encoder path (bypasses VAE)."""
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
vae_autocast_enabled = (
|
server_args, "vae", precision_attr="vae_precision"
|
||||||
vae_dtype != torch.float32
|
)
|
||||||
) and not server_args.disable_autocast
|
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
|
||||||
|
|
||||||
with torch.autocast(
|
with torch.autocast(
|
||||||
device_type=current_platform.device_type,
|
device_type=current_platform.device_type,
|
||||||
@@ -816,7 +853,9 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
def component_uses(
|
def component_uses(
|
||||||
self, server_args: ServerArgs, stage_name: str | None = None
|
self, server_args: ServerArgs, stage_name: str | None = None
|
||||||
) -> list[ComponentUse]:
|
) -> list[ComponentUse]:
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
|
server_args, self.component_name, precision_attr="vae_precision"
|
||||||
|
)
|
||||||
stage_name = self._component_stage_name(stage_name)
|
stage_name = self._component_stage_name(stage_name)
|
||||||
return [
|
return [
|
||||||
ComponentUse(
|
ComponentUse(
|
||||||
@@ -851,11 +890,11 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
server_args.pipeline_config, "prepare_condition_image_latent_ids", None
|
server_args.pipeline_config, "prepare_condition_image_latent_ids", None
|
||||||
)
|
)
|
||||||
condition_latents = [] if callable(prepare_condition_image_latent_ids) else None
|
condition_latents = [] if callable(prepare_condition_image_latent_ids) else None
|
||||||
# Setup VAE precision
|
# Setup VAE precision from user policy.
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
vae_autocast_enabled = (
|
server_args, self.component_name, precision_attr="vae_precision"
|
||||||
vae_dtype != torch.float32
|
)
|
||||||
) and not server_args.disable_autocast
|
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
|
||||||
|
|
||||||
with self.use_declared_component(
|
with self.use_declared_component(
|
||||||
component_name=self.component_name,
|
component_name=self.component_name,
|
||||||
@@ -902,14 +941,18 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
self.vae.enable_tiling()
|
self.vae.enable_tiling()
|
||||||
# if server_args.vae_sp:
|
# if server_args.vae_sp:
|
||||||
# self.vae.enable_parallel()
|
# self.vae.enable_parallel()
|
||||||
|
should_cast_vae = not vae_autocast_enabled
|
||||||
if not vae_autocast_enabled:
|
if not vae_autocast_enabled:
|
||||||
video_condition = video_condition.to(vae_dtype)
|
video_condition = video_condition.to(vae_dtype)
|
||||||
video_condition = server_args.pipeline_config.preprocess_vae_encode(
|
video_condition = server_args.pipeline_config.preprocess_vae_encode(
|
||||||
video_condition, self.vae
|
video_condition, self.vae
|
||||||
)
|
)
|
||||||
latent_dist: DiagonalGaussianDistribution = self.vae.encode(
|
with temporary_module_dtype(
|
||||||
video_condition
|
self.vae, vae_dtype, enabled=should_cast_vae
|
||||||
)
|
) as vae:
|
||||||
|
latent_dist: DiagonalGaussianDistribution = vae.encode(
|
||||||
|
video_condition
|
||||||
|
)
|
||||||
# for auto_encoder from diffusers
|
# for auto_encoder from diffusers
|
||||||
if isinstance(latent_dist, AutoencoderKLOutput):
|
if isinstance(latent_dist, AutoencoderKLOutput):
|
||||||
latent_dist = latent_dist.latent_dist
|
latent_dist = latent_dist.latent_dist
|
||||||
|
|||||||
+10
-5
@@ -24,6 +24,10 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||||
|
align_tensor_to_module_dtype,
|
||||||
|
get_module_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -733,7 +737,7 @@ class GlmImageBeforeDenoisingStage(PipelineStage):
|
|||||||
attention_kwargs = {}
|
attention_kwargs = {}
|
||||||
prompt_embeds = None
|
prompt_embeds = None
|
||||||
do_classifier_free_guidance = True
|
do_classifier_free_guidance = True
|
||||||
dtype = torch.bfloat16
|
dtype = get_module_dtype(self.transformer, torch.bfloat16)
|
||||||
|
|
||||||
self._guidance_scale = guidance_scale
|
self._guidance_scale = guidance_scale
|
||||||
self._current_timestep = None
|
self._current_timestep = None
|
||||||
@@ -799,14 +803,15 @@ class GlmImageBeforeDenoisingStage(PipelineStage):
|
|||||||
1, self.vae.config.latent_channels, 1, 1
|
1, self.vae.config.latent_channels, 1, 1
|
||||||
)
|
)
|
||||||
|
|
||||||
latents_mean = latents_mean.to(device=device, dtype=prompt_embeds.dtype)
|
vae_dtype = get_module_dtype(self.vae, prompt_embeds.dtype)
|
||||||
latents_std = latents_std.to(device=device, dtype=prompt_embeds.dtype)
|
latents_mean = latents_mean.to(device=device, dtype=vae_dtype)
|
||||||
|
latents_std = latents_std.to(device=device, dtype=vae_dtype)
|
||||||
|
|
||||||
for condition_image, condition_image_prior_token_id in zip(
|
for condition_image, condition_image_prior_token_id in zip(
|
||||||
ar_condition_images, prior_token_image_ids
|
ar_condition_images, prior_token_image_ids
|
||||||
):
|
):
|
||||||
condition_image = condition_image.to(
|
condition_image = align_tensor_to_module_dtype(
|
||||||
device=device, dtype=prompt_embeds.dtype
|
condition_image, self.vae, device=device
|
||||||
)
|
)
|
||||||
|
|
||||||
condition_latent = retrieve_latents(
|
condition_latent = retrieve_latents(
|
||||||
|
|||||||
+52
-25
@@ -9,7 +9,12 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import Decodin
|
|||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||||
|
align_tensor_to_module_dtype,
|
||||||
|
autocast_enabled,
|
||||||
|
resolve_precision,
|
||||||
|
temporary_module_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -32,9 +37,15 @@ class LTX2AVDecodingStage(DecodingStage):
|
|||||||
self, server_args: ServerArgs, stage_name: str | None = None
|
self, server_args: ServerArgs, stage_name: str | None = None
|
||||||
) -> list[ComponentUse]:
|
) -> list[ComponentUse]:
|
||||||
stage_name = self._component_stage_name(stage_name)
|
stage_name = self._component_stage_name(stage_name)
|
||||||
|
vae_dtype = resolve_precision(
|
||||||
|
server_args, "vae", precision_attr="vae_precision"
|
||||||
|
)
|
||||||
|
audio_vae_dtype = resolve_precision(
|
||||||
|
server_args, "audio_vae", precision_attr="audio_vae_precision"
|
||||||
|
)
|
||||||
return [
|
return [
|
||||||
ComponentUse(stage_name, "vae", target_dtype=torch.bfloat16),
|
ComponentUse(stage_name, "vae", target_dtype=vae_dtype),
|
||||||
ComponentUse(stage_name, "audio_vae"),
|
ComponentUse(stage_name, "audio_vae", target_dtype=audio_vae_dtype),
|
||||||
ComponentUse(stage_name, "vocoder"),
|
ComponentUse(stage_name, "vocoder"),
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -46,17 +57,18 @@ class LTX2AVDecodingStage(DecodingStage):
|
|||||||
def forward(self, batch: Req, server_args: ServerArgs) -> OutputBatch:
|
def forward(self, batch: Req, server_args: ServerArgs) -> OutputBatch:
|
||||||
self.load_model()
|
self.load_model()
|
||||||
|
|
||||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
vae_dtype = resolve_precision(
|
||||||
vae_autocast_enabled = (
|
server_args,
|
||||||
vae_dtype != torch.float32
|
"vae",
|
||||||
) and not server_args.disable_autocast
|
precision_attr="vae_precision",
|
||||||
|
)
|
||||||
|
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
|
||||||
|
|
||||||
original_dtype = vae_dtype
|
|
||||||
with self.use_declared_component(component_name="vae", module=self.vae) as vae:
|
with self.use_declared_component(component_name="vae", module=self.vae) as vae:
|
||||||
assert vae is not None
|
assert vae is not None
|
||||||
self.vae = vae
|
self.vae = vae
|
||||||
self.vae.eval()
|
self.vae.eval()
|
||||||
latents = batch.latents.to(get_local_torch_device(), dtype=torch.bfloat16)
|
latents = batch.latents.to(get_local_torch_device())
|
||||||
if self._ltx2_should_externally_denorm_video_latents(server_args):
|
if self._ltx2_should_externally_denorm_video_latents(server_args):
|
||||||
std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latents)
|
std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latents)
|
||||||
mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latents)
|
mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latents)
|
||||||
@@ -75,15 +87,19 @@ class LTX2AVDecodingStage(DecodingStage):
|
|||||||
self.vae.enable_tiling()
|
self.vae.enable_tiling()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
decode_output = self.vae.decode(latents)
|
should_cast_vae = not vae_autocast_enabled
|
||||||
|
if not vae_autocast_enabled:
|
||||||
|
latents = latents.to(vae_dtype)
|
||||||
|
with temporary_module_dtype(
|
||||||
|
self.vae, vae_dtype, enabled=should_cast_vae
|
||||||
|
) as vae:
|
||||||
|
decode_output = vae.decode(latents)
|
||||||
if isinstance(decode_output, tuple):
|
if isinstance(decode_output, tuple):
|
||||||
video = decode_output[0]
|
video = decode_output[0]
|
||||||
elif hasattr(decode_output, "sample"):
|
elif hasattr(decode_output, "sample"):
|
||||||
video = decode_output.sample
|
video = decode_output.sample
|
||||||
else:
|
else:
|
||||||
video = decode_output
|
video = decode_output
|
||||||
|
|
||||||
self.vae.to(original_dtype)
|
|
||||||
video = self.video_processor.postprocess_video(video, output_type="np")
|
video = self.video_processor.postprocess_video(video, output_type="np")
|
||||||
|
|
||||||
output_batch = OutputBatch(
|
output_batch = OutputBatch(
|
||||||
@@ -109,15 +125,12 @@ class LTX2AVDecodingStage(DecodingStage):
|
|||||||
assert audio_vae is not None
|
assert audio_vae is not None
|
||||||
self.audio_vae = audio_vae
|
self.audio_vae = audio_vae
|
||||||
self.audio_vae.eval()
|
self.audio_vae.eval()
|
||||||
try:
|
audio_vae_dtype = resolve_precision(
|
||||||
dtype = self.audio_vae.dtype
|
server_args,
|
||||||
except AttributeError:
|
"audio_vae",
|
||||||
dtype = None
|
precision_attr="audio_vae_precision",
|
||||||
if dtype is None:
|
)
|
||||||
try:
|
dtype = audio_vae_dtype
|
||||||
dtype = next(self.audio_vae.parameters()).dtype
|
|
||||||
except StopIteration:
|
|
||||||
dtype = torch.float32
|
|
||||||
audio_latents = audio_latents.to(device, dtype=dtype)
|
audio_latents = audio_latents.to(device, dtype=dtype)
|
||||||
try:
|
try:
|
||||||
latents_std = self.audio_vae.latents_std
|
latents_std = self.audio_vae.latents_std
|
||||||
@@ -147,11 +160,24 @@ class LTX2AVDecodingStage(DecodingStage):
|
|||||||
)
|
)
|
||||||
audio_latents = audio_latents * latents_std + latents_mean
|
audio_latents = audio_latents * latents_std + latents_mean
|
||||||
|
|
||||||
with torch.no_grad():
|
audio_vae_autocast_enabled = autocast_enabled(
|
||||||
|
audio_vae_dtype, server_args.disable_autocast
|
||||||
|
)
|
||||||
|
should_cast_audio_vae = not audio_vae_autocast_enabled
|
||||||
|
with torch.no_grad(), torch.autocast(
|
||||||
|
device_type=current_platform.device_type,
|
||||||
|
dtype=audio_vae_dtype,
|
||||||
|
enabled=audio_vae_autocast_enabled,
|
||||||
|
):
|
||||||
# Decode latents to spectrogram
|
# Decode latents to spectrogram
|
||||||
spectrogram = self.audio_vae.decode(
|
with temporary_module_dtype(
|
||||||
audio_latents, return_dict=False
|
self.audio_vae,
|
||||||
)[0]
|
audio_vae_dtype,
|
||||||
|
enabled=should_cast_audio_vae,
|
||||||
|
) as audio_vae:
|
||||||
|
spectrogram = audio_vae.decode(
|
||||||
|
audio_latents, return_dict=False
|
||||||
|
)[0]
|
||||||
|
|
||||||
with self.use_declared_component(
|
with self.use_declared_component(
|
||||||
component_name="vocoder",
|
component_name="vocoder",
|
||||||
@@ -170,6 +196,7 @@ class LTX2AVDecodingStage(DecodingStage):
|
|||||||
f"Vocoder expects channels*mel_bins={expected_in}, got {actual_in} from spectrogram shape {tuple(spectrogram.shape)}"
|
f"Vocoder expects channels*mel_bins={expected_in}, got {actual_in} from spectrogram shape {tuple(spectrogram.shape)}"
|
||||||
)
|
)
|
||||||
# Decode spectrogram to waveform
|
# Decode spectrogram to waveform
|
||||||
|
spectrogram = align_tensor_to_module_dtype(spectrogram, self.vocoder)
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
waveform = self.vocoder(spectrogram)
|
waveform = self.vocoder(spectrogram)
|
||||||
output_batch.audio = waveform.cpu().float()
|
output_batch.audio = waveform.cpu().float()
|
||||||
|
|||||||
+14
-2
@@ -20,6 +20,7 @@ from sglang.multimodal_gen.runtime.server_args import (
|
|||||||
is_ltx2_two_stage_pipeline_name,
|
is_ltx2_two_stage_pipeline_name,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -66,11 +67,22 @@ class LTX2AVLatentPreparationStage(LatentPreparationStage):
|
|||||||
):
|
):
|
||||||
if is_ltx23_native_variant(server_args.pipeline_config.vae_config.arch_config):
|
if is_ltx23_native_variant(server_args.pipeline_config.vae_config.arch_config):
|
||||||
if is_ltx2_two_stage_pipeline_name(server_args.pipeline_class_name):
|
if is_ltx2_two_stage_pipeline_name(server_args.pipeline_class_name):
|
||||||
return server_args.pipeline_config.get_latent_dtype(
|
prompt_dtype = (
|
||||||
batch.prompt_embeds[0].dtype
|
batch.prompt_embeds[0].dtype
|
||||||
|
if isinstance(batch.prompt_embeds, list)
|
||||||
|
else batch.prompt_embeds.dtype
|
||||||
)
|
)
|
||||||
|
return server_args.pipeline_config.get_latent_dtype(prompt_dtype)
|
||||||
|
if isinstance(batch.prompt_embeds, list) and batch.prompt_embeds:
|
||||||
|
return batch.prompt_embeds[0].dtype
|
||||||
|
if isinstance(batch.prompt_embeds, torch.Tensor):
|
||||||
|
return batch.prompt_embeds.dtype
|
||||||
return torch.float32
|
return torch.float32
|
||||||
return torch.float32
|
if isinstance(batch.prompt_embeds, list) and batch.prompt_embeds:
|
||||||
|
return batch.prompt_embeds[0].dtype
|
||||||
|
if isinstance(batch.prompt_embeds, torch.Tensor):
|
||||||
|
return batch.prompt_embeds.dtype
|
||||||
|
return resolve_precision(server_args, "dit", precision_attr="dit_precision")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _packed_video_latent_shape(
|
def _packed_video_latent_shape(
|
||||||
|
|||||||
+4
-2
@@ -17,6 +17,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
|||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.runtime.utils.precision import align_tensor_to_module_dtype
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -368,6 +369,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
|
|||||||
with self.use_declared_component(component_name="vae", module=self.vae) as vae:
|
with self.use_declared_component(component_name="vae", module=self.vae) as vae:
|
||||||
assert vae is not None
|
assert vae is not None
|
||||||
self.vae = vae
|
self.vae = vae
|
||||||
|
image = align_tensor_to_module_dtype(image, self.vae)
|
||||||
if isinstance(generator, list):
|
if isinstance(generator, list):
|
||||||
image_latents = [
|
image_latents = [
|
||||||
retrieve_latents(
|
retrieve_latents(
|
||||||
@@ -422,7 +424,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
|
|||||||
|
|
||||||
image_latents = None
|
image_latents = None
|
||||||
if image is not None:
|
if image is not None:
|
||||||
image = image.to(device=device, dtype=dtype)
|
image = align_tensor_to_module_dtype(image, self.vae, device=device)
|
||||||
if image.shape[1] != self.latent_channels:
|
if image.shape[1] != self.latent_channels:
|
||||||
image_latents = self._encode_vae_image(image=image, generator=generator)
|
image_latents = self._encode_vae_image(image=image, generator=generator)
|
||||||
else:
|
else:
|
||||||
@@ -510,7 +512,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
|
|||||||
image, calculated_height, calculated_width
|
image, calculated_height, calculated_width
|
||||||
)
|
)
|
||||||
image = image.unsqueeze(2)
|
image = image.unsqueeze(2)
|
||||||
image = image.to(dtype=self.vae_dtype)
|
image = align_tensor_to_module_dtype(image, self.vae, device=device)
|
||||||
|
|
||||||
prompt = batch.prompt
|
prompt = batch.prompt
|
||||||
with self.use_declared_component(
|
with self.use_declared_component(
|
||||||
|
|||||||
@@ -0,0 +1,119 @@
|
|||||||
|
from contextlib import contextmanager
|
||||||
|
from typing import Iterator, Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||||
|
|
||||||
|
|
||||||
|
def precision_to_dtype(precision: str, field_name: str = "precision") -> torch.dtype:
|
||||||
|
try:
|
||||||
|
return PRECISION_TO_TYPE[precision]
|
||||||
|
except KeyError as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"Unsupported {field_name}={precision!r}; "
|
||||||
|
f"expected one of {sorted(PRECISION_TO_TYPE)}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_precision(
|
||||||
|
server_args,
|
||||||
|
component_or_precision_attr: str,
|
||||||
|
*,
|
||||||
|
precision_attr: Optional[str] = None,
|
||||||
|
field_name: Optional[str] = None,
|
||||||
|
) -> torch.dtype:
|
||||||
|
precision_attr = precision_attr or component_or_precision_attr
|
||||||
|
precision = getattr(server_args.pipeline_config, precision_attr)
|
||||||
|
return precision_to_dtype(precision, field_name or precision_attr)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_component_precision(server_args, module_name: str) -> Optional[torch.dtype]:
|
||||||
|
pipeline_config = getattr(server_args, "pipeline_config", None)
|
||||||
|
if pipeline_config is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if module_name in ("audio_vae", "vocoder"):
|
||||||
|
precision_attr = "audio_vae_precision"
|
||||||
|
elif module_name in ("vae", "video_vae"):
|
||||||
|
precision_attr = "vae_precision"
|
||||||
|
elif module_name in (
|
||||||
|
"transformer",
|
||||||
|
"transformer_2",
|
||||||
|
"audio_dit",
|
||||||
|
"video_dit",
|
||||||
|
"connectors",
|
||||||
|
"dual_tower_bridge",
|
||||||
|
):
|
||||||
|
precision_attr = "dit_precision"
|
||||||
|
elif module_name == "image_encoder":
|
||||||
|
precision_attr = "image_encoder_precision"
|
||||||
|
elif module_name == "text_encoder" or module_name.startswith("text_encoder_"):
|
||||||
|
precisions = getattr(pipeline_config, "text_encoder_precisions", None)
|
||||||
|
if not precisions:
|
||||||
|
return None
|
||||||
|
suffix = module_name.removeprefix("text_encoder")
|
||||||
|
index = 0 if suffix == "" else int(suffix.removeprefix("_")) - 1
|
||||||
|
if index < 0 or index >= len(precisions):
|
||||||
|
raise ValueError(
|
||||||
|
f"No configured precision for {module_name!r}; "
|
||||||
|
f"text_encoder_precisions has {len(precisions)} entries"
|
||||||
|
)
|
||||||
|
precision = precisions[index]
|
||||||
|
return precision_to_dtype(precision, f"text_encoder_precisions[{index}]")
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if not hasattr(pipeline_config, precision_attr):
|
||||||
|
return None
|
||||||
|
return resolve_precision(server_args, precision_attr)
|
||||||
|
|
||||||
|
|
||||||
|
def autocast_enabled(dtype: torch.dtype, disable_autocast: bool) -> bool:
|
||||||
|
return dtype != torch.float32 and not disable_autocast
|
||||||
|
|
||||||
|
|
||||||
|
def get_module_dtype(module, default: torch.dtype = torch.float32) -> torch.dtype:
|
||||||
|
try:
|
||||||
|
return next(module.parameters()).dtype
|
||||||
|
except (AttributeError, StopIteration):
|
||||||
|
dtype = getattr(module, "dtype", None)
|
||||||
|
return dtype if isinstance(dtype, torch.dtype) else default
|
||||||
|
|
||||||
|
|
||||||
|
def align_tensor_to_module_dtype(
|
||||||
|
tensor: torch.Tensor,
|
||||||
|
module,
|
||||||
|
*,
|
||||||
|
device: Optional[Union[torch.device, str]] = None,
|
||||||
|
default_dtype: torch.dtype = torch.float32,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
dtype = get_module_dtype(module, default=default_dtype)
|
||||||
|
if device is None:
|
||||||
|
try:
|
||||||
|
device = next(module.parameters()).device
|
||||||
|
except (AttributeError, StopIteration):
|
||||||
|
device = tensor.device
|
||||||
|
if not tensor.is_floating_point():
|
||||||
|
return tensor.to(device=device)
|
||||||
|
return tensor.to(device=device, dtype=dtype)
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def temporary_module_dtype(
|
||||||
|
module,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
*,
|
||||||
|
enabled: bool = True,
|
||||||
|
restore_dtype: Optional[torch.dtype] = None,
|
||||||
|
) -> Iterator:
|
||||||
|
if not enabled:
|
||||||
|
yield module
|
||||||
|
return
|
||||||
|
|
||||||
|
original_dtype = restore_dtype or get_module_dtype(module)
|
||||||
|
module = module.to(dtype=dtype)
|
||||||
|
try:
|
||||||
|
yield module
|
||||||
|
finally:
|
||||||
|
module.to(dtype=original_dtype)
|
||||||
@@ -0,0 +1,173 @@
|
|||||||
|
import importlib.util
|
||||||
|
import sys
|
||||||
|
import types
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
def _load_precision_module():
|
||||||
|
stub_names = (
|
||||||
|
"sglang",
|
||||||
|
"sglang.multimodal_gen",
|
||||||
|
"sglang.multimodal_gen.runtime",
|
||||||
|
"sglang.multimodal_gen.runtime.utils",
|
||||||
|
"sglang.multimodal_gen.utils",
|
||||||
|
)
|
||||||
|
missing = object()
|
||||||
|
previous_modules = {name: sys.modules.get(name, missing) for name in stub_names}
|
||||||
|
|
||||||
|
try:
|
||||||
|
utils_module = types.ModuleType("sglang.multimodal_gen.utils")
|
||||||
|
utils_module.PRECISION_TO_TYPE = {
|
||||||
|
"fp16": torch.float16,
|
||||||
|
"bf16": torch.bfloat16,
|
||||||
|
"fp32": torch.float32,
|
||||||
|
}
|
||||||
|
for package_name in stub_names[:-1]:
|
||||||
|
package = types.ModuleType(package_name)
|
||||||
|
package.__path__ = []
|
||||||
|
sys.modules[package_name] = package
|
||||||
|
sys.modules["sglang.multimodal_gen.utils"] = utils_module
|
||||||
|
|
||||||
|
precision_path = (
|
||||||
|
Path(__file__).resolve().parents[2] / "runtime/utils/precision.py"
|
||||||
|
)
|
||||||
|
spec = importlib.util.spec_from_file_location(
|
||||||
|
"_diffusion_precision_under_test", precision_path
|
||||||
|
)
|
||||||
|
precision = importlib.util.module_from_spec(spec)
|
||||||
|
sys.modules[spec.name] = precision
|
||||||
|
spec.loader.exec_module(precision)
|
||||||
|
finally:
|
||||||
|
for module_name, previous_module in previous_modules.items():
|
||||||
|
if previous_module is missing:
|
||||||
|
sys.modules.pop(module_name, None)
|
||||||
|
else:
|
||||||
|
sys.modules[module_name] = previous_module
|
||||||
|
|
||||||
|
return precision
|
||||||
|
|
||||||
|
|
||||||
|
precision = _load_precision_module()
|
||||||
|
align_tensor_to_module_dtype = precision.align_tensor_to_module_dtype
|
||||||
|
autocast_enabled = precision.autocast_enabled
|
||||||
|
get_module_dtype = precision.get_module_dtype
|
||||||
|
precision_to_dtype = precision.precision_to_dtype
|
||||||
|
resolve_component_precision = precision.resolve_component_precision
|
||||||
|
resolve_precision = precision.resolve_precision
|
||||||
|
temporary_module_dtype = precision.temporary_module_dtype
|
||||||
|
|
||||||
|
|
||||||
|
class _DtypedNoParameterModule(torch.nn.Module):
|
||||||
|
def __init__(self, dtype: torch.dtype):
|
||||||
|
super().__init__()
|
||||||
|
self.dtype = dtype
|
||||||
|
|
||||||
|
|
||||||
|
class _ParameterDtypeWinsModule(torch.nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.dtype = torch.float32
|
||||||
|
self.weight = torch.nn.Parameter(torch.ones(1, dtype=torch.float16))
|
||||||
|
|
||||||
|
|
||||||
|
class TestDiffusionPrecisionConsistency(unittest.TestCase):
|
||||||
|
def _server_args(self, **overrides):
|
||||||
|
config = {
|
||||||
|
"vae_precision": "fp16",
|
||||||
|
"audio_vae_precision": "bf16",
|
||||||
|
"dit_precision": "fp32",
|
||||||
|
"image_encoder_precision": "fp16",
|
||||||
|
"text_encoder_precisions": ["fp16", "bf16"],
|
||||||
|
}
|
||||||
|
config.update(overrides)
|
||||||
|
return SimpleNamespace(pipeline_config=SimpleNamespace(**config))
|
||||||
|
|
||||||
|
def test_precision_lookup(self):
|
||||||
|
server_args = self._server_args()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
resolve_precision(server_args, "vae", precision_attr="vae_precision"),
|
||||||
|
torch.float16,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
resolve_precision(server_args, "dit", precision_attr="dit_precision"),
|
||||||
|
torch.float32,
|
||||||
|
)
|
||||||
|
with self.assertRaisesRegex(ValueError, "Unsupported vae_precision"):
|
||||||
|
resolve_precision(self._server_args(vae_precision="fp8"), "vae_precision")
|
||||||
|
with self.assertRaisesRegex(ValueError, "Unsupported custom_precision"):
|
||||||
|
precision_to_dtype("fp8", "custom_precision")
|
||||||
|
|
||||||
|
def test_component_precision_mapping(self):
|
||||||
|
server_args = self._server_args()
|
||||||
|
expected = {
|
||||||
|
"vae": torch.float16,
|
||||||
|
"video_vae": torch.float16,
|
||||||
|
"audio_vae": torch.bfloat16,
|
||||||
|
"vocoder": torch.bfloat16,
|
||||||
|
"transformer": torch.float32,
|
||||||
|
"transformer_2": torch.float32,
|
||||||
|
"audio_dit": torch.float32,
|
||||||
|
"video_dit": torch.float32,
|
||||||
|
"connectors": torch.float32,
|
||||||
|
"dual_tower_bridge": torch.float32,
|
||||||
|
"image_encoder": torch.float16,
|
||||||
|
"text_encoder": torch.float16,
|
||||||
|
"text_encoder_2": torch.bfloat16,
|
||||||
|
}
|
||||||
|
|
||||||
|
for module_name, expected_dtype in expected.items():
|
||||||
|
self.assertEqual(
|
||||||
|
resolve_component_precision(server_args, module_name),
|
||||||
|
expected_dtype,
|
||||||
|
module_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsNone(resolve_component_precision(SimpleNamespace(), "vae"))
|
||||||
|
self.assertIsNone(
|
||||||
|
resolve_component_precision(server_args, "unregistered_component")
|
||||||
|
)
|
||||||
|
self.assertIsNone(
|
||||||
|
resolve_component_precision(
|
||||||
|
self._server_args(text_encoder_precisions=[]), "text_encoder"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_autocast_and_dtype_alignment(self):
|
||||||
|
self.assertTrue(autocast_enabled(torch.float16, disable_autocast=False))
|
||||||
|
self.assertTrue(autocast_enabled(torch.bfloat16, disable_autocast=False))
|
||||||
|
self.assertFalse(autocast_enabled(torch.float32, disable_autocast=False))
|
||||||
|
self.assertFalse(autocast_enabled(torch.float16, disable_autocast=True))
|
||||||
|
|
||||||
|
module = _ParameterDtypeWinsModule()
|
||||||
|
self.assertEqual(get_module_dtype(module), torch.float16)
|
||||||
|
aligned = align_tensor_to_module_dtype(
|
||||||
|
torch.ones(1, dtype=torch.float32), module
|
||||||
|
)
|
||||||
|
self.assertEqual(aligned.dtype, torch.float16)
|
||||||
|
|
||||||
|
module_without_parameters = _DtypedNoParameterModule(torch.bfloat16)
|
||||||
|
self.assertEqual(get_module_dtype(module_without_parameters), torch.bfloat16)
|
||||||
|
tokens = torch.ones(2, dtype=torch.long)
|
||||||
|
aligned_tokens = align_tensor_to_module_dtype(tokens, module_without_parameters)
|
||||||
|
self.assertEqual(aligned_tokens.dtype, torch.long)
|
||||||
|
|
||||||
|
def test_temporary_module_dtype(self):
|
||||||
|
module = torch.nn.Linear(2, 2).to(dtype=torch.float32)
|
||||||
|
|
||||||
|
with temporary_module_dtype(module, torch.bfloat16):
|
||||||
|
self.assertEqual(module.weight.dtype, torch.bfloat16)
|
||||||
|
|
||||||
|
self.assertEqual(module.weight.dtype, torch.float32)
|
||||||
|
|
||||||
|
with temporary_module_dtype(module, torch.float16, enabled=False) as casted:
|
||||||
|
self.assertIs(casted, module)
|
||||||
|
self.assertEqual(module.weight.dtype, torch.float32)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user