[diffusion] chore: honor explicit component offload (#36931)

This commit is contained in:
Mick
2026-08-29 14:29:57 +08:00
committed by GitHub
parent a25df83fe3
commit f5319af1c9
2 changed files with 31 additions and 2 deletions
@@ -1540,11 +1540,11 @@ class ServerArgs(DisaggServerArgsMixin):
def require_component_resident(
self, component_name: str, *, feature_name: str
) -> None:
configured_mode = self.canonical_residency_mode(component_name)
configured_mode = self.explicit_residency_mode(component_name)
if configured_mode is not None and configured_mode != RESIDENT:
raise ValueError(
f"{feature_name} requires {component_name!r} to be resident; "
f"got {configured_mode!r} from --component-residency"
f"got {configured_mode!r} from an explicit residency option"
)
self._required_resident_components.add(component_name)
@@ -1233,6 +1233,35 @@ class TestOffloadDefaults(unittest.TestCase):
args.disable_fsdp_for_component("text_encoder")
self.assertFalse(args.should_use_fsdp_for_component("text_encoder"))
def test_resident_requirement_rejects_every_explicit_offload_surface(self):
cases = (
{"component_residency": ["text_encoder=component-offload"]},
{"component_residency": ["text_encoder=layerwise-offload"]},
{"cpu_offload_components": ["text_encoder"]},
{"text_encoder_cpu_offload": True},
{"layerwise_offload_components": ["text_encoder"]},
)
for kwargs in cases:
with self.subTest(kwargs=kwargs):
args = self._from_dict_with_task_type(
ModelTaskType.T2V,
kwargs={"performance_mode": "manual", **kwargs},
)
with self.assertRaisesRegex(ValueError, "explicit residency option"):
args.require_component_resident(
"text_encoder", feature_name="test backend"
)
def test_resident_requirement_can_override_automatic_placement(self):
args = self._from_dict_with_task_type(ModelTaskType.T2V, memory_gb=16)
self.assertIsNone(args.explicit_residency_mode("text_encoder"))
self.assertNotEqual(args.residency_mode("text_encoder"), RESIDENT)
args.require_component_resident("text_encoder", feature_name="test backend")
self.assertEqual(args.residency_mode("text_encoder"), RESIDENT)
def test_diffusers_component_residency_is_pipeline_wide(self):
self.assertFalse(resolve_diffusers_pipeline_offload({"all": RESIDENT}))
self.assertTrue(resolve_diffusers_pipeline_offload({"all": COMPONENT_OFFLOAD}))