Text Classification
Transformers
Joblib
English
cybersecurity
industrial-control-systems
bert
from-scratch
synthetic-data
Instructions to use ARotting/protocol-guardian with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use ARotting/protocol-guardian with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="ARotting/protocol-guardian")# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("ARotting/protocol-guardian", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| from __future__ import annotations | |
| import json | |
| import random | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| import trackio | |
| from model import build_model, parameter_count | |
| from sklearn.metrics import ( | |
| accuracy_score, | |
| confusion_matrix, | |
| f1_score, | |
| precision_score, | |
| recall_score, | |
| ) | |
| from torch.nn import functional as F | |
| from torch.utils.data import DataLoader, TensorDataset | |
| from transformers import PreTrainedTokenizerFast | |
| PROJECT_DIR = Path(__file__).resolve().parent | |
| ROOT_DIR = PROJECT_DIR.parents[1] | |
| TOKENIZER_DIR = ROOT_DIR / "projects" / "snip-0.4m" / "artifacts" / "snip-0.4m-base" | |
| DATA_DIR = PROJECT_DIR / "data" | |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "protocol-guardian-bert" | |
| def seed_everything(seed: int) -> None: | |
| random.seed(seed) | |
| np.random.seed(seed) | |
| torch.manual_seed(seed) | |
| def load_split(name: str, tokenizer) -> TensorDataset: | |
| frame = pd.read_parquet(DATA_DIR / f"{name}.parquet") | |
| encoded = tokenizer( | |
| frame["text"].tolist(), | |
| padding="max_length", | |
| truncation=True, | |
| max_length=64, | |
| return_tensors="pt", | |
| add_special_tokens=True, | |
| ) | |
| return TensorDataset( | |
| encoded["input_ids"], | |
| encoded["attention_mask"], | |
| torch.tensor(frame["label"].to_numpy(dtype=np.int64, copy=True)), | |
| ) | |
| def evaluate(model, dataset: TensorDataset) -> dict: | |
| model.eval() | |
| loader = DataLoader(dataset, batch_size=256, shuffle=False) | |
| labels, predictions, losses = [], [], [] | |
| for input_ids, attention_mask, targets in loader: | |
| logits = model( | |
| input_ids=input_ids, | |
| attention_mask=attention_mask, | |
| ).logits | |
| losses.append(F.cross_entropy(logits, targets, reduction="sum").item()) | |
| labels.extend(targets.tolist()) | |
| predictions.extend(logits.argmax(dim=1).tolist()) | |
| return { | |
| "loss": float(sum(losses) / len(labels)), | |
| "accuracy": float(accuracy_score(labels, predictions)), | |
| "precision": float(precision_score(labels, predictions)), | |
| "recall": float(recall_score(labels, predictions)), | |
| "f1": float(f1_score(labels, predictions)), | |
| "confusion_matrix": confusion_matrix(labels, predictions).tolist(), | |
| "examples": len(labels), | |
| } | |
| def main() -> None: | |
| seed_everything(2026) | |
| tokenizer = PreTrainedTokenizerFast.from_pretrained(TOKENIZER_DIR) | |
| train_dataset = load_split("train", tokenizer) | |
| validation_dataset = load_split("validation", tokenizer) | |
| test_dataset = load_split("test", tokenizer) | |
| model = build_model(len(tokenizer), tokenizer.pad_token_id) | |
| loader = DataLoader( | |
| train_dataset, | |
| batch_size=96, | |
| shuffle=True, | |
| generator=torch.Generator().manual_seed(2026), | |
| ) | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=0.0015, weight_decay=0.01) | |
| epochs = 10 | |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs) | |
| best_accuracy = -1.0 | |
| best_epoch = 0 | |
| best_state = None | |
| trackio.init( | |
| project="protocol-guardian", | |
| name="tiny-bert-from-scratch-v1", | |
| config={ | |
| "parameters": parameter_count(model), | |
| "train_examples": len(train_dataset), | |
| "validation_examples": len(validation_dataset), | |
| "held_out_template_examples": len(test_dataset), | |
| "epochs": epochs, | |
| }, | |
| ) | |
| for epoch in range(1, epochs + 1): | |
| model.train() | |
| running_loss = 0.0 | |
| examples = 0 | |
| for input_ids, attention_mask, labels in loader: | |
| logits = model( | |
| input_ids=input_ids, | |
| attention_mask=attention_mask, | |
| ).logits | |
| loss = F.cross_entropy(logits, labels) | |
| optimizer.zero_grad(set_to_none=True) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| optimizer.step() | |
| running_loss += loss.item() * len(labels) | |
| examples += len(labels) | |
| scheduler.step() | |
| validation = evaluate(model, validation_dataset) | |
| trackio.log( | |
| { | |
| "epoch": epoch, | |
| "train_loss": running_loss / examples, | |
| "validation_loss": validation["loss"], | |
| "validation_accuracy": validation["accuracy"], | |
| "validation_f1": validation["f1"], | |
| "learning_rate": scheduler.get_last_lr()[0], | |
| } | |
| ) | |
| if validation["accuracy"] > best_accuracy: | |
| best_accuracy = validation["accuracy"] | |
| best_epoch = epoch | |
| best_state = { | |
| key: value.detach().cpu().clone() | |
| for key, value in model.state_dict().items() | |
| } | |
| trackio.finish() | |
| if best_state is None: | |
| sys.exit("Training did not produce a checkpoint.") | |
| model.load_state_dict(best_state) | |
| test = evaluate(model, test_dataset) | |
| ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) | |
| model.save_pretrained(ARTIFACT_DIR, safe_serialization=True) | |
| tokenizer.save_pretrained(ARTIFACT_DIR) | |
| summary = { | |
| "model": "Protocol Guardian Tiny BERT", | |
| "parameters": parameter_count(model), | |
| "best_epoch": best_epoch, | |
| "best_validation_accuracy": best_accuracy, | |
| "test_split": "1200 examples from entirely held-out command templates", | |
| "test": test, | |
| "limitations": [ | |
| "Synthetic English command corpus", | |
| "Binary research classifier only", | |
| "Not a replacement for deterministic authorization and safety logic", | |
| ], | |
| } | |
| (ARTIFACT_DIR / "training_summary.json").write_text( | |
| json.dumps(summary, indent=2), | |
| encoding="utf-8", | |
| ) | |
| print(json.dumps(summary, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |