[diffusion] Reject unsafe quality=high BCG replay (#36008)

This commit is contained in:
Xiaoyu Zhang
2026-08-24 08:50:26 +08:00
committed by GitHub
parent f4448e677f
commit 447048dba2
4 changed files with 61 additions and 1 deletions
@@ -659,6 +659,19 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
mounted_fusions.add(description)
else:
unmount(transformer)
if want and mounted_fusions and self.server_args.enable_breakable_cuda_graph:
for transformer in filter(None, [self.transformer, self.transformer_2]):
for _, _, unmount in _QUALITY_FUSION_HANDLERS:
unmount(transformer)
descriptions = ", ".join(sorted(mounted_fusions))
raise ValueError(
"quality='high' cannot be used with breakable CUDA graphs for "
f"this model because its request-scoped DiT fusions "
f"({descriptions}) do not match the lossless warmup graphs. "
"Disable breakable CUDA graphs or use quality='lossless'."
)
self._quality_fusions_mounted = want
for description in sorted(mounted_fusions):
logger.info("Mounted %s for quality=high", description)
@@ -58,6 +58,44 @@ class SanaVideoTransformer3DModel(torch.nn.Module):
pass
class TestQualityFusionBCGCompatibility(unittest.TestCase):
def setUp(self):
self.stage = DenoisingStage.__new__(DenoisingStage)
self.stage.server_args = SimpleNamespace(enable_breakable_cuda_graph=True)
self.stage.transformer = OtherTransformer2DModel()
self.stage.transformer_2 = None
self.stage._quality_fusions_mounted = False
@staticmethod
def _batch(quality: str):
return SimpleNamespace(sampling_params=SimpleNamespace(quality=quality))
def test_rejects_high_when_dit_fusion_would_replace_captured_graph(self):
unmounted = []
handlers = (
(
"test fusion",
lambda _: True,
lambda transformer: unmounted.append(transformer),
),
)
with patch.object(denoising_module, "_QUALITY_FUSION_HANDLERS", handlers):
with self.assertRaisesRegex(ValueError, "lossless warmup graphs"):
self.stage._maybe_toggle_quality_fusions(self._batch("high"))
self.assertEqual(unmounted, [self.stage.transformer])
self.assertFalse(self.stage._quality_fusions_mounted)
def test_allows_high_when_model_has_no_dit_quality_fusions(self):
handlers = (("test fusion", lambda _: False, lambda _: None),)
with patch.object(denoising_module, "_QUALITY_FUSION_HANDLERS", handlers):
self.stage._maybe_toggle_quality_fusions(self._batch("high"))
self.assertTrue(self.stage._quality_fusions_mounted)
def _fake_cache_dit_batch(*, is_warmup: bool) -> SimpleNamespace:
return SimpleNamespace(
is_warmup=is_warmup,