Download server.py from Cortiqa/guard-a40M: direct link, hf CLI and curl.
- Browser
- Download file 3.16 kB
-
https://huggingface.co/Cortiqa/guard-a40M/resolve/main/server.py
- Command line
-
hf download hf://Cortiqa/guard-a40M/server.py
-
curl -L -o server.py https://huggingface.co/Cortiqa/guard-a40M/resolve/main/server.py
3.16 kB
| """ | |
| Cortiqa-Guard-40M Production HTTP Microservice (FastAPI) | |
| Deploy anywhere: Docker, Cloud (AWS/GCP/Azure), VPS, Kubernetes, or Local Server. | |
| """ | |
| import time | |
| from pathlib import Path | |
| from typing import Dict, List, Optional | |
| import torch | |
| from fastapi import FastAPI, HTTPException | |
| from pydantic import BaseModel | |
| from transformers import AutoTokenizer | |
| from model import CortiqaGuard40M, load_config, GUARD_LABELS | |
| app = FastAPI( | |
| title="Cortiqa-Guard-40M Universal AI Firewall", | |
| description="High-Speed (<3ms) Prompt Injection & Safety Guardrail API", | |
| version="1.0.0", | |
| ) | |
| # Global model container | |
| guard_model: Optional[CortiqaGuard40M] = None | |
| guard_tokenizer = None | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| class ScanRequest(BaseModel): | |
| prompt: str | |
| max_length: Optional[int] = 512 | |
| class ScanResponse(BaseModel): | |
| is_safe: bool | |
| action: str # "ALLOW" or "BLOCK" | |
| threat_category: str | |
| class_id: int | |
| confidence: float | |
| latency_ms: float | |
| scores: Dict[str, float] | |
| def load_firewall(): | |
| global guard_model, guard_tokenizer | |
| cfg_path = Path(__file__).parent / "configs" / "guard_40m.yaml" | |
| cfg = load_config(str(cfg_path)) | |
| print(f"Loading tokenizer: {cfg.training.tokenizer_name}...") | |
| guard_tokenizer = AutoTokenizer.from_pretrained(cfg.training.tokenizer_name) | |
| if guard_tokenizer.pad_token is None: | |
| guard_tokenizer.pad_token = guard_tokenizer.eos_token | |
| print(f"Loading Cortiqa Guard 40M on {device}...") | |
| guard_model = CortiqaGuard40M(cfg.model).to(device) | |
| # Check for trained checkpoint | |
| ckpt_path = Path(__file__).parent / "checkpoints" / "guard_final.pt" | |
| if ckpt_path.exists(): | |
| print(f"Loading trained weights from {ckpt_path}...") | |
| ckpt = torch.load(ckpt_path, map_location=device) | |
| guard_model.load_state_dict(ckpt["model_state"]) | |
| guard_model.eval() | |
| print("✅ Cortiqa Guard Firewall API is live and ready to protect LLMs!") | |
| def health(): | |
| return { | |
| "status": "healthy", | |
| "device": str(device), | |
| "parameters": guard_model.param_count() if guard_model else 0, | |
| "firewall_active": True, | |
| } | |
| def scan_prompt(req: ScanRequest): | |
| if not req.prompt.strip(): | |
| raise HTTPException(status_code=400, detail="Empty prompt provided.") | |
| t0 = time.perf_counter() | |
| enc = guard_tokenizer( | |
| req.prompt, | |
| return_tensors="pt", | |
| max_length=req.max_length, | |
| truncation=True, | |
| ).to(device) | |
| with torch.no_grad(): | |
| res = guard_model.inspect_prompt(enc["input_ids"]) | |
| latency = (time.perf_counter() - t0) * 1000.0 | |
| return ScanResponse( | |
| is_safe=res["is_safe"], | |
| action="ALLOW" if res["is_safe"] else "BLOCK", | |
| threat_category=res["label"], | |
| class_id=res["class_id"], | |
| confidence=res["confidence"], | |
| latency_ms=round(latency, 2), | |
| scores=res["all_scores"], | |
| ) | |
| if __name__ == "__main__": | |
| import uvicorn | |
| uvicorn.run(app, host="0.0.0.0", port=8000) | |