diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py index c1f92703e..dbefbc9b4 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py @@ -38,6 +38,10 @@ from sglang.multimodal_gen.runtime.utils.precision import ( resolve_component_precision, resolve_decode_precision, ) +from sglang.multimodal_gen.runtime.weights.source import ( + materialize_weight, + resolve_weight, +) from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.srt.model_loader.checkpoint_quantization import ( resolve_checkpoint_quant_spec, @@ -272,6 +276,25 @@ class VAELoader(ComponentLoader): component_names = ["vae", "audio_vae", "video_vae"] expected_library = "diffusers" + @staticmethod + def resolve_model_weights_path( + component_model_path: str, + server_args: ServerArgs, + component_name: str, + ) -> str: + weights_override = getattr(server_args, "component_weights_paths", {}).get( + component_name + ) + if weights_override is None: + return component_model_path + model_weights_path = materialize_weight(resolve_weight(weights_override)) + logger.info( + "Using weight-file override for %s: %s", + component_name, + model_weights_path, + ) + return model_weights_path + def customized_load_kwargs_for_component( self, server_args: ServerArgs, component_name: str ) -> dict[str, bool]: @@ -295,6 +318,11 @@ class VAELoader(ComponentLoader): cpu_offload_flag: bool = False, ): """Load the VAE based on the model path, and inference args.""" + component_weights_path = self.resolve_model_weights_path( + component_model_path, + server_args, + component_name, + ) config = get_diffusers_component_config(component_path=component_model_path) server_args.model_paths[component_name] = component_model_path native_only = component_name in getattr( @@ -340,6 +368,11 @@ class VAELoader(ComponentLoader): auto_map = config.get("auto_map", {}) auto_model_map = auto_map.get("AutoModel") + if auto_model_map and component_weights_path != component_model_path: + raise ComponentCheckpointUnsupportedError( + f"{component_name!r} uses a custom Diffusers class that cannot " + "consume a weights-only override" + ) if auto_model_map and not native_only: module_path, cls_name = auto_model_map.rsplit(".", 1) custom_module_file = os.path.join(component_model_path, f"{module_path}.py") @@ -374,17 +407,25 @@ class VAELoader(ComponentLoader): vae_cls, _ = ModelRegistry.resolve_model_cls(class_name) vae = vae_cls(vae_config).to(target_device) - safetensors_list = _list_safetensors_files(component_model_path) - safetensors_list = server_args.pipeline_config.select_vae_weight_files( - safetensors_list=safetensors_list, - component_model_path=component_model_path, - component_name=component_name, - vae_precision=vae_precision, - ) + if os.path.isfile(component_weights_path): + if not component_weights_path.endswith(".safetensors"): + raise ValueError( + f"VAE weight overrides must be safetensors, got " + f"{component_weights_path!r}" + ) + safetensors_list = [component_weights_path] + else: + safetensors_list = _list_safetensors_files(component_weights_path) + safetensors_list = server_args.pipeline_config.select_vae_weight_files( + safetensors_list=safetensors_list, + component_model_path=component_weights_path, + component_name=component_name, + vae_precision=vae_precision, + ) assert ( len(safetensors_list) >= 1 - ), f"Found no safetensors files in {component_model_path}" + ), f"Found no safetensors files in {component_weights_path}" loaded = {} for sf_path in safetensors_list: loaded.update(safetensors_load_file(sf_path)) @@ -437,7 +478,7 @@ class VAELoader(ComponentLoader): logger.info("VAE: converted %d Conv3d weights to channels_last_3d", n) _hold_decoder_weights_in_decode_dtype( - vae, server_args, component_name, component_model_path + vae, server_args, component_name, component_weights_path ) vae = current_platform.optimize_vae(vae) return vae diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index 654430f43..29fc41f1c 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -1705,6 +1705,7 @@ class ServerArgs(DisaggServerArgsMixin): is_dit_component_name(component) or is_text_encoder_component_name(component) or is_image_encoder_component_name(component) + or is_vae_component_name(component) ) if ( not supports_weight_file_override diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index e8dc31a0d..bce674ba3 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -170,6 +170,7 @@ class TestServerArgsPathExpansion(unittest.TestCase): "model_path": "/data/my-model", "component_paths": { "text_encoder": "owner/repo/text_encoder/model.safetensors", + "audio_vae": "owner/repo/vae/audio.safetensors", "vae": "owner/repo/vae", }, } @@ -178,7 +179,10 @@ class TestServerArgsPathExpansion(unittest.TestCase): self.assertEqual(args.component_paths, {"vae": "owner/repo/vae"}) self.assertEqual( args.component_weights_paths, - {"text_encoder": "owner/repo/text_encoder/model.safetensors"}, + { + "text_encoder": "owner/repo/text_encoder/model.safetensors", + "audio_vae": "owner/repo/vae/audio.safetensors", + }, ) def test_supplemental_weight_file_remains_a_component_path(self): diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py index 02927db63..eadacd8f3 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py @@ -123,6 +123,28 @@ class TestMatchCheckpointDtypes(unittest.TestCase): class TestVAELoader(unittest.TestCase): + def test_weights_override_keeps_base_component_config(self): + loader = vae_loader.VAELoader() + server_args = _FakeServerArgs(QwenImagePipelineConfig()) + server_args.component_weights_paths = { + "audio_vae": "owner/repo/audio_vae.safetensors" + } + + with ( + patch.object(vae_loader, "resolve_weight", return_value="resolved"), + patch.object( + vae_loader, + "materialize_weight", + return_value="/cache/audio.safetensors", + ), + ): + self.assertEqual( + loader.resolve_model_weights_path( + "/base/audio_vae", server_args, "audio_vae" + ), + "/cache/audio.safetensors", + ) + def test_mps_layerwise_load_uses_residency_api(self): loader = vae_loader.VAELoader() server_args = _FakeServerArgs(QwenImagePipelineConfig())