[diffusion] feat: support vae weight-file overrides (#36085)
This commit is contained in:
@@ -38,6 +38,10 @@ from sglang.multimodal_gen.runtime.utils.precision import (
|
|||||||
resolve_component_precision,
|
resolve_component_precision,
|
||||||
resolve_decode_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.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||||
from sglang.srt.model_loader.checkpoint_quantization import (
|
from sglang.srt.model_loader.checkpoint_quantization import (
|
||||||
resolve_checkpoint_quant_spec,
|
resolve_checkpoint_quant_spec,
|
||||||
@@ -272,6 +276,25 @@ class VAELoader(ComponentLoader):
|
|||||||
component_names = ["vae", "audio_vae", "video_vae"]
|
component_names = ["vae", "audio_vae", "video_vae"]
|
||||||
expected_library = "diffusers"
|
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(
|
def customized_load_kwargs_for_component(
|
||||||
self, server_args: ServerArgs, component_name: str
|
self, server_args: ServerArgs, component_name: str
|
||||||
) -> dict[str, bool]:
|
) -> dict[str, bool]:
|
||||||
@@ -295,6 +318,11 @@ class VAELoader(ComponentLoader):
|
|||||||
cpu_offload_flag: bool = False,
|
cpu_offload_flag: bool = False,
|
||||||
):
|
):
|
||||||
"""Load the VAE based on the model path, and inference args."""
|
"""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)
|
config = get_diffusers_component_config(component_path=component_model_path)
|
||||||
server_args.model_paths[component_name] = component_model_path
|
server_args.model_paths[component_name] = component_model_path
|
||||||
native_only = component_name in getattr(
|
native_only = component_name in getattr(
|
||||||
@@ -340,6 +368,11 @@ class VAELoader(ComponentLoader):
|
|||||||
|
|
||||||
auto_map = config.get("auto_map", {})
|
auto_map = config.get("auto_map", {})
|
||||||
auto_model_map = auto_map.get("AutoModel")
|
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:
|
if auto_model_map and not native_only:
|
||||||
module_path, cls_name = auto_model_map.rsplit(".", 1)
|
module_path, cls_name = auto_model_map.rsplit(".", 1)
|
||||||
custom_module_file = os.path.join(component_model_path, f"{module_path}.py")
|
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_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||||
vae = vae_cls(vae_config).to(target_device)
|
vae = vae_cls(vae_config).to(target_device)
|
||||||
|
|
||||||
safetensors_list = _list_safetensors_files(component_model_path)
|
if os.path.isfile(component_weights_path):
|
||||||
safetensors_list = server_args.pipeline_config.select_vae_weight_files(
|
if not component_weights_path.endswith(".safetensors"):
|
||||||
safetensors_list=safetensors_list,
|
raise ValueError(
|
||||||
component_model_path=component_model_path,
|
f"VAE weight overrides must be safetensors, got "
|
||||||
component_name=component_name,
|
f"{component_weights_path!r}"
|
||||||
vae_precision=vae_precision,
|
)
|
||||||
)
|
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 (
|
assert (
|
||||||
len(safetensors_list) >= 1
|
len(safetensors_list) >= 1
|
||||||
), f"Found no safetensors files in {component_model_path}"
|
), f"Found no safetensors files in {component_weights_path}"
|
||||||
loaded = {}
|
loaded = {}
|
||||||
for sf_path in safetensors_list:
|
for sf_path in safetensors_list:
|
||||||
loaded.update(safetensors_load_file(sf_path))
|
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)
|
logger.info("VAE: converted %d Conv3d weights to channels_last_3d", n)
|
||||||
|
|
||||||
_hold_decoder_weights_in_decode_dtype(
|
_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)
|
vae = current_platform.optimize_vae(vae)
|
||||||
return vae
|
return vae
|
||||||
|
|||||||
@@ -1705,6 +1705,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
is_dit_component_name(component)
|
is_dit_component_name(component)
|
||||||
or is_text_encoder_component_name(component)
|
or is_text_encoder_component_name(component)
|
||||||
or is_image_encoder_component_name(component)
|
or is_image_encoder_component_name(component)
|
||||||
|
or is_vae_component_name(component)
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
not supports_weight_file_override
|
not supports_weight_file_override
|
||||||
|
|||||||
@@ -170,6 +170,7 @@ class TestServerArgsPathExpansion(unittest.TestCase):
|
|||||||
"model_path": "/data/my-model",
|
"model_path": "/data/my-model",
|
||||||
"component_paths": {
|
"component_paths": {
|
||||||
"text_encoder": "owner/repo/text_encoder/model.safetensors",
|
"text_encoder": "owner/repo/text_encoder/model.safetensors",
|
||||||
|
"audio_vae": "owner/repo/vae/audio.safetensors",
|
||||||
"vae": "owner/repo/vae",
|
"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_paths, {"vae": "owner/repo/vae"})
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
args.component_weights_paths,
|
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):
|
def test_supplemental_weight_file_remains_a_component_path(self):
|
||||||
|
|||||||
@@ -123,6 +123,28 @@ class TestMatchCheckpointDtypes(unittest.TestCase):
|
|||||||
|
|
||||||
|
|
||||||
class TestVAELoader(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):
|
def test_mps_layerwise_load_uses_residency_api(self):
|
||||||
loader = vae_loader.VAELoader()
|
loader = vae_loader.VAELoader()
|
||||||
server_args = _FakeServerArgs(QwenImagePipelineConfig())
|
server_args = _FakeServerArgs(QwenImagePipelineConfig())
|
||||||
|
|||||||
Reference in New Issue
Block a user