memguard / tests /test_mitigation.py
asthanarohan's picture
Mitigation
9d3899c
Raw
History Blame Contribute Delete
5.71 kB
"""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