- Удалены поле maxChars, параметр конструктора и метод evictWhileOverLimit. - Хранилище теперь ограничено только TTL (ttl-minutes), без вытеснения по объёму. - Конструкторы переведены на (int ttlMinutes) и (int ttlMinutes, SharedIndex, PayloadCipher). - Обновлены тесты и PipelineWarmup на новые сигнатуры.
158 lines
5.3 KiB
Java
158 lines
5.3 KiB
Java
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 для адресов.
|
|
*
|
|
* <p>Набор {@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<BenchmarkFixtures.Sample> samples = BenchmarkFixtures.load("/benchmark-two-model.txt");
|
|
Map<String, Score> byType = new LinkedHashMap<>();
|
|
Map<String, List<String>> missed = new LinkedHashMap<>();
|
|
|
|
for (BenchmarkFixtures.Sample sample : samples) {
|
|
List<Span> 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<String, Score> 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<Span> 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<String, Score> 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<String, Score> 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<String, Score> 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<String, List<String>> 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);
|
|
}
|
|
}
|