zeekalph/devra-ai_verifier
0
1from fastapi import FastAPI, File, Form, UploadFile2from pydantic import BaseModel3from typing import List, Dict4import torch5from transformers import AutoTokenizer, AutoModelForMaskedLM6from sentence_transformers import SentenceTransformer, util7import torchvision.models as models8from torchvision import transforms9from PIL import Image10import io11import zipfile12import pandas as pd13import numpy as np14import gc15import os16 17os.environ["TOKENIZERS_PARALLELISM"] = "false"18 19app = FastAPI(title="AI Dataset Verifier")20 21device = torch.device("cpu")22 23# Lazy load24_tokenizer = None25_model = None26_sentence_model = None27_resnet = None28_transform = None29 30def get_tokenizer():31 global _tokenizer32 if _tokenizer is None:33 _tokenizer = AutoTokenizer.from_pretrained("prajjwal1/bert-tiny")34 return _tokenizer35 36def get_model():37 global _model38 if _model is None:39 _model = AutoModelForMaskedLM.from_pretrained("prajjwal1/bert-tiny").to(device)40 _model.eval()41 return _model42 43def get_sentence_model():44 global _sentence_model45 if _sentence_model is None:46 _sentence_model = SentenceTransformer('all-MiniLM-L6-v2')47 return _sentence_model48 49def get_resnet():50 global _resnet, _transform51 if _resnet is None:52 _resnet = models.resnet18(pretrained=True).to(device)53 _resnet.eval()54 _transform = transforms.Compose([55 transforms.Resize(256),56 transforms.CenterCrop(224),57 transforms.ToTensor(),58 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])59 ])60 return _resnet, _transform61 62class Response(BaseModel):63 scores: Dict[str, int]64 status: str65 issues: List[str] = []66 67def score_text(texts: List[str], desc: str = None):68 if not texts:69 return {"quality": 0, "completeness": 0, "consistency": 0, "relevance": 50}70 tokenizer = get_tokenizer()71 model = get_model()72 perps = []73 for t in texts[:2]:74 enc = tokenizer(t, return_tensors="pt", truncation=True, max_length=128).to(device)75 with torch.no_grad():76 loss = model(**enc, labels=enc["input_ids"]).loss77 perps.append(torch.exp(loss).item())78 del enc; gc.collect()79 quality = max(0, min(100, 100 - np.mean(perps) * 2))80 relevance = 5081 if desc and texts:82 sm = get_sentence_model()83 e1 = sm.encode(desc, convert_to_tensor=True)84 e2 = sm.encode(texts[:3], convert_to_tensor=True)85 sim = util.cos_sim(e1, e2).mean().item()86 relevance = int((sim + 1) * 50)87 return {88 "quality": int(quality),89 "completeness": 100 if len(texts) >= 2 else 50,90 "consistency": 90,91 "relevance": relevance92 }93 94@app.post("/verify", response_model=Response)95async def verify(file: UploadFile = File(...), description: str = Form(None)):96 content = await file.read()97 texts = []98 try:99 with zipfile.ZipFile(io.BytesIO(content)) as z:100 for n in z.namelist():101 if n.endswith(('.csv', '.txt')):102 data = z.read(n)103 if n.endswith('.csv'):104 df = pd.read_csv(io.BytesIO(data))105 texts.extend(df.astype(str).values.flatten().tolist())106 else:107 texts.append(data.decode())108 except:109 try:110 texts = [content.decode()]111 except:112 pass113 scores = score_text(texts, description)114 status = "VERIFIED" if scores["quality"] >= 60 else "FAILED"115 return Response(scores=scores, status=status, issues=[])116 117@app.get("/")118def home():119 return {"message": "AI Verifier API – Use /docs"}