Download train_classifier.py from dougdotcon/fiscalcheck-ptbr-compatibility-classifier: direct link, hf CLI and curl.
- Browser
- Download file 6.95 kB
-
https://huggingface.co/dougdotcon/fiscalcheck-ptbr-compatibility-classifier/resolve/main/train_classifier.py
- Command line
-
hf download hf://dougdotcon/fiscalcheck-ptbr-compatibility-classifier/train_classifier.py
-
curl -L -o train_classifier.py https://huggingface.co/dougdotcon/fiscalcheck-ptbr-compatibility-classifier/resolve/main/train_classifier.py
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() | |