Spaces:
Sleeping
Sleeping
| """Mitigation tests. | |
| These exercise the real mitigation code path (text-encoder dropout averaging + | |
| the detect -> mitigate -> block decision) against a tiny toy pipeline -- no | |
| Stable Diffusion download. torch is required, so the module is skipped where | |
| torch is absent (e.g. the minimal CI install); the memguard env has torch, so it | |
| runs there. | |
| The toy UNet emits ``mean(tanh(embedding)) * base``, so the memorization signal | |
| ``||eps_c - eps_uc||`` is monotonic in the gap between the conditional and | |
| unconditional embeddings. The toy text encoder models attention dropout as a | |
| pull of the conditional embedding toward the unconditional one, so averaging | |
| dropout draws (the mitigation) reduces that gap. | |
| """ | |
| import types | |
| import pytest | |
| torch = pytest.importorskip("torch") | |
| from memguard.detector import MemorizationDetector # noqa: E402 | |
| from memguard.mitigation import mitigate_prompt_embeddings # noqa: E402 | |
| from memguard.pipeline import GuardedDiffusionPipeline # noqa: E402 | |
| C, S, SEQ, DIM = 4, 8, 4, 6 | |
| SD1_CAL = (1.04, 0.96) # so the toy score spans the 0.9 threshold like SD1.4 | |
| class _Unet: | |
| def __init__(self): | |
| self.config = types.SimpleNamespace(in_channels=C, sample_size=S) | |
| self.dtype = torch.float32 | |
| self.base = torch.ones(1, C, S, S) | |
| def __call__(self, x, t, encoder_hidden_states=None, return_dict=True): | |
| a = torch.tanh(encoder_hidden_states).mean(dim=(1, 2)) # [B] | |
| out = a.view(-1, 1, 1, 1) * self.base | |
| return (out,) if not return_dict else types.SimpleNamespace(sample=out) | |
| class _Scheduler: | |
| def __init__(self): | |
| self.alphas_cumprod = torch.linspace(0.9999, 0.001, 1000) | |
| self.init_noise_sigma = 1.0 | |
| self.timesteps = torch.arange(49, -1, -1) | |
| def set_timesteps(self, n, device=None): | |
| self.timesteps = torch.arange(n - 1, -1, -1) | |
| def scale_model_input(self, x, t): | |
| return x | |
| class _TextEncoder(torch.nn.Module): | |
| """Minimal CLIP-like encoder exposing a float ``dropout`` attribute.""" | |
| def __init__(self): | |
| super().__init__() | |
| self.dropout = 0.0 | |
| class _Pipe: | |
| """Minimal stand-in implementing just the surface mitigation/metric use.""" | |
| def __init__(self, cond_val=0.5, uncond_val=0.1): | |
| self.unet = _Unet() | |
| self.scheduler = _Scheduler() | |
| self.text_encoder = _TextEncoder().eval() # inference pipelines keep this in eval | |
| self._execution_device = "cpu" | |
| self._cond_val = cond_val | |
| self._uncond_val = uncond_val | |
| self.calls = [] | |
| def encode_prompt(self, prompt, device, num_images_per_prompt=1, | |
| do_classifier_free_guidance=True, negative_prompt=None): | |
| te = self.text_encoder | |
| cv = self._cond_val | |
| if te.training and te.dropout > 0: | |
| # Attention dropout pulls the conditional embedding toward the | |
| # unconditional one (breaking the memorization trigger), with noise. | |
| frac = min(1.0, te.dropout * 3.0) # p=0.3 -> ~0.9 pull | |
| cv = (self._cond_val - (self._cond_val - self._uncond_val) * frac | |
| + float(torch.randn(1)) * 0.02) | |
| cond = torch.full((1, SEQ, DIM), cv) | |
| uncond = torch.full((1, SEQ, DIM), self._uncond_val) | |
| return cond, uncond | |
| def __call__(self, prompt=None, *, prompt_embeds=None, negative_prompt_embeds=None, | |
| num_inference_steps=50, guidance_scale=7.5, latents=None, **kwargs): | |
| self.calls.append("embeds" if prompt_embeds is not None else "prompt") | |
| return types.SimpleNamespace(images=[object()]) | |
| def test_dropout_mitigation_reduces_signal_and_restores_state(): | |
| torch.manual_seed(0) | |
| pipe = _Pipe(cond_val=0.5, uncond_val=0.1) | |
| latents = torch.ones(1, C, S, S) | |
| out = mitigate_prompt_embeddings( | |
| pipe, "memorized prompt", latents=latents, | |
| num_inference_steps=50, num_samples=10, dropout_p=0.3, | |
| threshold=0.9, calibration=SD1_CAL, | |
| ) | |
| assert out["samples"] == 10 | |
| assert out["signal_after"] < out["signal_before"] | |
| assert out["passed"] is True | |
| assert torch.isfinite(out["prompt_embeds"]).all() | |
| # the text encoder is returned to its original state (dropout off, eval mode) | |
| assert pipe.text_encoder.dropout == 0.0 | |
| assert pipe.text_encoder.training is False | |
| def test_guarded_mitigates_when_possible(): | |
| torch.manual_seed(0) | |
| pipe = _Pipe(cond_val=0.5, uncond_val=0.1) # strong signal -> memorized | |
| guarded = GuardedDiffusionPipeline( | |
| pipe, MemorizationDetector(threshold=0.9, calibration=SD1_CAL), mitigate=True, | |
| ) | |
| res = guarded.generate("memorized prompt", num_inference_steps=5, seed=0) | |
| assert res["mitigated"] is True | |
| assert res["blocked"] is False | |
| assert res["memorized"] is False | |
| assert res["image"] is not None | |
| assert res["signal_after"] < res["signal_before"] | |
| def test_guarded_blocks_without_mitigation(): | |
| pipe = _Pipe(cond_val=0.5, uncond_val=0.1) | |
| guarded = GuardedDiffusionPipeline( | |
| pipe, MemorizationDetector(threshold=0.9), mitigate=False, | |
| ) | |
| res = guarded.generate("memorized prompt", num_inference_steps=5, seed=0) | |
| assert res["memorized"] is True | |
| assert res["blocked"] is True | |
| assert res["image"] is None | |
| assert res["mitigated"] is False | |
| def test_guarded_passes_benign(): | |
| pipe = _Pipe(cond_val=0.12, uncond_val=0.1) # tiny gap -> not memorized | |
| guarded = GuardedDiffusionPipeline(pipe, MemorizationDetector(threshold=0.9)) | |
| res = guarded.generate("benign prompt", num_inference_steps=5, seed=0) | |
| assert res["memorized"] is False | |
| assert res["blocked"] is False | |
| assert res["mitigated"] is False | |
| assert res["image"] is not None | |