memguard / tests /test_detector.py
asthanarohan's picture
Fix
1751b09
Raw
History Blame Contribute Delete
2.34 kB
"""Detector tests.
The detector scores a ``(pipe, prompt)`` pair with the real anisotropy metric
(``prompt -> Stable Diffusion -> memorization score``). CI has no GPU/model, so
we inject a fake ``score_fn`` to exercise the threshold/decision logic and the
config plumbing without downloading Stable Diffusion. The metric itself is
covered end-to-end by tests/test_metrics.py (its pure numpy core) and
examples/score_prompt_sd1.py (a live model run).
"""
from memguard.detector import MemorizationDetector
_PIPE = object() # opaque stand-in; the fake score_fn never touches it
def _fake_score_fn(expected=0.95):
"""A score_fn that flags any prompt containing 'mem', echoing kwargs used."""
calls = {}
def score_fn(pipe, prompt, **kwargs):
calls["pipe"] = pipe
calls["prompt"] = prompt
calls["kwargs"] = kwargs
return expected if "mem" in prompt.lower() else 0.05
score_fn.calls = calls
return score_fn
def test_score_in_range():
det = MemorizationDetector(score_fn=_fake_score_fn())
assert 0.0 <= det.score(_PIPE, "a benign prompt") <= 1.0
def test_is_memorized_returns_bool():
det = MemorizationDetector(threshold=0.9, score_fn=_fake_score_fn())
assert det.is_memorized(_PIPE, "a memorized prompt") is True
assert det.is_memorized(_PIPE, "a benign prompt") is False
def test_check_has_expected_keys():
det = MemorizationDetector(score_fn=_fake_score_fn())
out = det.check(_PIPE, "a memorized prompt")
assert set(out) == {"prompt", "score", "memorized", "threshold"}
assert out["threshold"] == 0.9
assert out["prompt"] == "a memorized prompt"
assert out["memorized"] is True
def test_threshold_controls_decision():
fn = _fake_score_fn(expected=0.5)
assert MemorizationDetector(threshold=0.0, score_fn=fn).is_memorized(_PIPE, "mem") is True
assert MemorizationDetector(threshold=1.0001, score_fn=fn).is_memorized(_PIPE, "mem") is False
def test_detector_passes_config_to_metric():
fn = _fake_score_fn()
det = MemorizationDetector(sd_ver=2, num_inference_steps=25, score_fn=fn)
det.score(_PIPE, "mem", latents="LAT")
assert fn.calls["pipe"] is _PIPE
assert fn.calls["kwargs"]["sd_ver"] == 2
assert fn.calls["kwargs"]["num_inference_steps"] == 25
assert fn.calls["kwargs"]["latents"] == "LAT"