package ru.pdguard; import static org.junit.jupiter.api.Assertions.assertTrue; import java.util.ArrayList; import java.util.Comparator; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Optional; import org.junit.jupiter.api.Test; import ru.pdguard.config.SystemPolicy; import ru.pdguard.core.PayloadStore; import ru.pdguard.core.Pipeline; import ru.pdguard.detect.NameCascade; import ru.pdguard.detect.RuleRegistry; import ru.pdguard.detect.Span; import ru.pdguard.mask.Masker; /** * Оценка двухмодельной архитектуры: WikiNEuRal для имён, ruBERT для адресов. * *

Набор {@code benchmark-two-model.txt} проверяет, что имена клиентов и адреса маскируются, а * известные личности — нет. Для каждого типа считается посимвольная точность, полнота и F1. */ class TwoModelBenchmarkTest { private final Pipeline pipeline = new Pipeline( new RuleRegistry(), new Masker(), new PayloadStore(30), new NameCascade( new NameCascade.EngineConfig( "wikineural", Optional.of("models/wikineural-ner"), "rubert", Optional.of("models/rubert-ner"), "off", Optional.empty()), 16, 4)); private static final class Score { private int truePositive; private int falsePositive; private int falseNegative; private int gold() { return truePositive + falseNegative; } private double precision() { int found = truePositive + falsePositive; return found == 0 ? 1.0 : (double) truePositive / found; } private double recall() { return gold() == 0 ? 1.0 : (double) truePositive / gold(); } private double f1() { double p = precision(); double r = recall(); return p + r == 0 ? 0.0 : 2 * p * r / (p + r); } } @Test void twoModelEfficiency() { List samples = BenchmarkFixtures.load("/benchmark-two-model.txt"); Map byType = new LinkedHashMap<>(); Map> missed = new LinkedHashMap<>(); for (BenchmarkFixtures.Sample sample : samples) { List found = pipeline.findPersonalData(sample.text(), SystemPolicy.DEFAULT); String[] goldChars = paint(sample.text().length(), sample.gold()); String[] foundChars = paint(sample.text().length(), found); for (int i = 0; i < sample.text().length(); i++) { account(byType, goldChars[i], foundChars[i]); } for (Span gold : sample.gold()) { boolean hit = found.stream().anyMatch(f -> f.type().equals(gold.type()) && f.overlaps(gold)); if (!hit) { missed .computeIfAbsent(gold.type(), t -> new ArrayList<>()) .add(sample.text().substring(gold.start(), gold.end())); } } } report(byType); reportMissed(missed); // Каждый тип должен быть найден с F1 не ниже 0.8. for (Map.Entry e : byType.entrySet()) { assertTrue( e.getValue().f1() >= 0.8, String.format("F1 по типу %s упал до %.3f", e.getKey(), e.getValue().f1())); } } private static String[] paint(int length, List spans) { String[] painted = new String[length]; for (Span span : spans) { for (int i = span.start(); i < Math.min(span.end(), length); i++) { painted[i] = span.type(); } } return painted; } private static void account(Map byType, String gold, String found) { if (gold != null) { Score score = byType.computeIfAbsent(gold, t -> new Score()); if (gold.equals(found)) { score.truePositive++; } else { score.falseNegative++; } } if (found != null && !found.equals(gold)) { byType.computeIfAbsent(found, t -> new Score()).falsePositive++; } } private void report(Map byType) { StringBuilder out = new StringBuilder(2048); out.append("\n=== Эффективность двухмодельной архитектуры ===\n\n"); out.append( String.format("%-20s %8s %8s %8s %8s%n", "тип", "знаков", "точность", "полнота", "F1")); byType.entrySet().stream() .sorted( Comparator.comparingInt((Map.Entry e) -> e.getValue().gold()).reversed()) .forEach( e -> out.append( String.format( "%-20s %8d %8.3f %8.3f %8.3f%n", e.getKey(), e.getValue().gold(), e.getValue().precision(), e.getValue().recall(), e.getValue().f1()))); System.out.println(out); } private void reportMissed(Map> missed) { if (missed.isEmpty()) { return; } StringBuilder out = new StringBuilder(); out.append("\n=== Не распознанные значения по типам ===\n"); missed.forEach( (type, values) -> out.append(type).append(": ").append(String.join(" | ", values)).append('\n')); System.out.println(out); } }