"""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()