[diffusion] fix: don't route an unreadable checkpoint into the native fallback (#39292)

Co-authored-by: Mick Qian <mickqian@users.noreply.github.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-09-13 20:30:24 +08:00
committed by GitHub
co-authored by Mick Qian Claude Opus 5
parent 14b647cf27
commit 9ffe548738
2 changed files with 39 additions and 0 deletions
@@ -397,6 +397,12 @@ class ComponentLoader(ABC):
ComponentAttentionBackendNotAppliedError,
ComponentCheckpointUnsupportedError,
ComponentResidencyError,
# the native fallback answers "there is no customized implementation
# for this architecture"; a checkpoint that cannot be read is a
# different failure and the fallback cannot read it either, so let
# the original error name the file instead of reporting it as a
# missing implementation
OSError,
):
raise
except Exception as e:
@@ -129,6 +129,39 @@ class TestTransformerLoaderFallbackAdmission(unittest.TestCase):
self._server_args(**overrides), "transformer_2"
)
def test_unreadable_checkpoint_is_not_a_missing_implementation(self):
# the native fallback answers "no customized implementation for this
# architecture"; a checkpoint that cannot be read is a different failure,
# and routing it into the fallback reports a missing implementation for a
# model that has one
loader = TransformerLoader()
missing_shard = FileNotFoundError(
2, "No such file or directory", "/cache/transformer/shard-00002.safetensors"
)
customized_load = mock.patch.object(
loader, "_load_customized_with_context", side_effect=missing_shard
)
native_load = mock.patch.object(
loader, "_load_native_with_context", return_value=object()
)
available_memory = mock.patch(
"sglang.multimodal_gen.runtime.loader.component_loaders."
"component_loader.current_platform.get_available_gpu_memory",
return_value=0.0,
)
with customized_load, native_load as native, available_memory:
with self.assertRaises(FileNotFoundError) as caught:
loader.load(
"/model/transformer_2",
self._server_args(),
"transformer_2",
"diffusers",
)
self.assertIn("shard-00002.safetensors", str(caught.exception))
native.assert_not_called()
def test_replicated_execution_keeps_native_fallback_available(self):
self.assertIsNone(
TransformerLoader().validate_native_fallback(