refactor: подтянуть Legal NER, переформатировать код и добавить тесты
Слияние с 685ec97 (третья ступень NER для юридических реквизитов),
код приведён к google-java-format, добавлены юнит-тесты
AdaptiveConcurrencyLimiter/SystemsConfig/PayloadCipher.
This commit is contained in:
@@ -1,5 +1,13 @@
|
||||
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;
|
||||
@@ -9,134 +17,141 @@ import ru.pdguard.detect.RuleRegistry;
|
||||
import ru.pdguard.detect.Span;
|
||||
import ru.pdguard.mask.Masker;
|
||||
|
||||
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 static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
|
||||
/**
|
||||
* Оценка двухмодельной архитектуры: WikiNEuRal для имён, ruBERT для адресов.
|
||||
*
|
||||
* <p>Набор {@code benchmark-two-model.txt} проверяет, что имена клиентов и адреса
|
||||
* маскируются, а известные личности — нет. Для каждого типа считается посимвольная
|
||||
* точность, полнота и F1.
|
||||
* <p>Набор {@code benchmark-two-model.txt} проверяет, что имена клиентов и адреса маскируются, а
|
||||
* известные личности — нет. Для каждого типа считается посимвольная точность, полнота и F1.
|
||||
*/
|
||||
class TwoModelBenchmarkTest {
|
||||
|
||||
private final Pipeline pipeline = new Pipeline(new RuleRegistry(), new Masker(),
|
||||
new PayloadStore(10_000_000L, 30),
|
||||
new NameCascade(
|
||||
new NameCascade.EngineConfig(
|
||||
"wikineural", Optional.of("models/wikineural-ner"),
|
||||
"rubert", Optional.of("models/rubert-ner"),
|
||||
"off", Optional.empty()),
|
||||
16, 4));
|
||||
private final Pipeline pipeline =
|
||||
new Pipeline(
|
||||
new RuleRegistry(),
|
||||
new Masker(),
|
||||
new PayloadStore(10_000_000L, 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 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);
|
||||
}
|
||||
private int gold() {
|
||||
return truePositive + falseNegative;
|
||||
}
|
||||
|
||||
@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 double precision() {
|
||||
int found = truePositive + falsePositive;
|
||||
return found == 0 ? 1.0 : (double) truePositive / found;
|
||||
}
|
||||
|
||||
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 double recall() {
|
||||
return gold() == 0 ? 1.0 : (double) truePositive / gold();
|
||||
}
|
||||
|
||||
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 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()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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);
|
||||
}
|
||||
report(byType);
|
||||
reportMissed(missed);
|
||||
|
||||
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);
|
||||
// Каждый тип должен быть найден с 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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user