"""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