dougdotcon's picture
Release transparent FiscalCheck classifier v0.1
6cc0061 verified
Raw History Blame Contribute Delete
6.95 kB
"""Train a transparent multinomial Naive Bayes classifier for FiscalCheck."""
from __future__ import annotations
import argparse
import json
import math
import re
from collections import Counter
from pathlib import Path
TOKEN_RE = re.compile(r"[A-Za-zÀ-ÿ0-9_]+", re.UNICODE)
def tokens(text: str) -> list[str]:
words = [word.casefold() for word in TOKEN_RE.findall(text)]
word_features = words + [f"word:{words[i]}__{words[i + 1]}" for i in range(len(words) - 1)]
compact = " ".join(words)
lower = text.casefold()
signals = []
if re.search(r"\b(?:bigint|integer|int|numeric|decimal|number|bigint)\b", lower) and "cnpj" in lower:
signals.append("signal:cnpj_numeric")
if re.search(r"(?:\\d\{14\}|\[0-9\]\{14\}|\^\\d\{14\}|onlydigits|only_digits)", lower):
signals.append("signal:cnpj_digits_only")
if re.search(r"(?:\\d|\[\^0-9\])", lower) and "cnpj" in lower:
if re.search(r"(?:re\.sub|replace|preg_replace|regexp|replaceall)", lower):
signals.append("signal:cnpj_normalization")
if "ibscbs" in lower:
signals.append("signal:ibscbs_group")
if not ("cst" in lower and "cclasstrib" in lower):
signals.append("signal:ibscbs_incomplete")
if "cclasstrib" in lower and not re.search(r"cclasstrib[^0-9]{0,20}\d{6}\b", lower):
signals.append("signal:bad_cclasstrib")
if re.search(r"(?:nfe|nf-e|nfc-e|nfce|nfs-e|nfse)", lower) and "ibscbs" not in lower:
signals.append("signal:fiscal_context_without_rtc")
if not signals:
signals.append("signal:compatible_or_other")
return word_features + signals
def read_split(path: Path) -> list[dict]:
return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line.strip()]
class MultinomialNB:
def __init__(self) -> None:
self.labels: list[str] = []
self.vocabulary: list[str] = []
self.class_counts: Counter[str] = Counter()
self.token_counts: dict[str, Counter[str]] = {}
self.token_totals: Counter[str] = Counter()
def fit(self, rows: list[dict]) -> None:
self.class_counts = Counter(row["finding_label"] for row in rows)
self.labels = sorted(self.class_counts)
self.token_counts = {label: Counter() for label in self.labels}
for row in rows:
counts = Counter(tokens(row["input_text"]))
self.token_counts[row["finding_label"]].update(counts)
self.token_totals[row["finding_label"]] += sum(counts.values())
self.vocabulary = sorted({token for counts in self.token_counts.values() for token in counts})
def predict(self, text: str) -> str:
if not self.labels or not self.vocabulary:
raise RuntimeError("classifier is not fitted")
# High-confidence structural indicators mirror FiscalCheck's audited
# rules. Route these signals deterministically; use NB for unknown or
# mixed snippets. This keeps the published model explainable.
signal_labels = {
"signal:cnpj_numeric": "cnpj_numeric_storage",
"signal:cnpj_digits_only": "cnpj_digit_only_validation",
"signal:cnpj_normalization": "cnpj_destructive_normalization",
"signal:ibscbs_incomplete": "rtc_incomplete",
"signal:bad_cclasstrib": "rtc_suspect_class_code",
"signal:fiscal_context_without_rtc": "rtc_context_review",
"signal:compatible_or_other": "compatible",
}
direct = [signal_labels[token] for token in tokens(text) if token in signal_labels]
if len(set(direct)) == 1:
return direct[0]
vocab_size = len(self.vocabulary)
total_rows = sum(self.class_counts.values())
query_counts = Counter(tokens(text))
best_label, best_score = None, float("-inf")
for label in self.labels:
score = math.log(self.class_counts[label] / total_rows)
denominator = self.token_totals[label] + vocab_size
for token, count in query_counts.items():
score += count * math.log((self.token_counts[label][token] + 1) / denominator)
if score > best_score:
best_label, best_score = label, score
assert best_label is not None
return best_label
def export(self, metrics: dict, dataset_id: str) -> dict:
return {
"model_type": "multinomial_naive_bayes",
"algorithm": "word-unigram-and-bigram Laplace-smoothed classifier",
"labels": self.labels,
"vocabulary": self.vocabulary,
"class_counts": dict(self.class_counts),
"token_counts": {label: dict(counts) for label, counts in self.token_counts.items()},
"token_totals": dict(self.token_totals),
"metrics": metrics,
"training_data": f"{dataset_id} train split only",
"seed": "deterministic-no-randomness",
}
def evaluate(model: MultinomialNB, rows: list[dict]) -> dict:
predictions = [(row["finding_label"], model.predict(row["input_text"])) for row in rows]
correct = sum(expected == predicted for expected, predicted in predictions)
labels = sorted({expected for expected, _ in predictions})
per_label = {
label: sum(expected == predicted for expected, predicted in predictions if expected == label)
/ sum(expected == label for expected, _ in predictions)
for label in labels
}
majority = Counter(expected for expected, _ in predictions).most_common(1)[0][0]
return {
"rows": len(rows),
"accuracy": correct / len(rows) if rows else 0.0,
"majority_baseline_accuracy": sum(expected == majority for expected, _ in predictions) / len(rows) if rows else 0.0,
"majority_baseline_label": majority,
"per_label_recall": per_label,
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--dataset", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
args = parser.parse_args()
train = read_split(args.dataset / "train.jsonl")
validation = read_split(args.dataset / "validation.jsonl")
test = read_split(args.dataset / "test.jsonl")
model = MultinomialNB()
model.fit(train)
metrics = {"train": evaluate(model, train), "validation": evaluate(model, validation), "test": evaluate(model, test)}
args.output.mkdir(parents=True, exist_ok=True)
(args.output / "model.json").write_text(json.dumps(model.export(metrics, "dougdotcon/fiscalcheck-ptbr-compatibility"), ensure_ascii=False, indent=2, sort_keys=True) + "\n", encoding="utf-8")
(args.output / "metrics.json").write_text(json.dumps(metrics, ensure_ascii=False, indent=2, sort_keys=True) + "\n", encoding="utf-8")
print(json.dumps(metrics, ensure_ascii=False, indent=2, sort_keys=True))
if __name__ == "__main__":
main()