Files
CNN_Website/api/server.py
T

156 lines
3.5 KiB
Python

from enum import Enum
import os
import sqlite3
from dotenv import load_dotenv
from pydantic import BaseModel
import torch
from torchvision import transforms
from fastapi import Depends, FastAPI, File, HTTPException, Header, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from PIL import Image
import io
import torch.nn.functional as F
from bird_cnn import Bird_CNN
import threading
BUILD_PATH = "./build_models"
IMAGE_SIZE = 64
sem = threading.Semaphore(1) #adjust to performance
class bird_species(Enum):
Common_Kingfisher = 0
Common_Myna = 1
House_Crow = 2
Indian_Peacock = 3
Indian_Pitta = 4
Ruddy_Shelduck = 5
Sarus_Crane = 6
transform = transforms.Compose([
transforms.Resize(IMAGE_SIZE),
transforms.CenterCrop(IMAGE_SIZE),
transforms.ToTensor()
])
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = Bird_CNN(c_in=3, c_hidden=16, c_out=7)
full_path = os.path.join(BUILD_PATH, "bird_cnn")
model.load_state_dict(torch.load(full_path, map_location=torch.device(device)))
model.to(device)
model.eval()
app = FastAPI()
origins = [
"https://marvinkrausser.com",
"https://api.marvinkrausser.com",
]
app.add_middleware(
CORSMiddleware,
allow_origins=origins,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.post("/predict")
async def predict(file: UploadFile = File(...)):
with sem:
image_bytes = await file.read()
image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
image = transform(image).unsqueeze(0).to(device)
with torch.no_grad():
pred = model(image)
probs = F.softmax(pred, dim=1)
confidence, cls = torch.max(probs, dim=1)
return {
"class": bird_species(cls.item()).name,
"confidence": confidence.item()
}
load_dotenv("database/.env")
API_KEY = os.getenv("API_KEY")
def get_api_key(authorization: str = Header(None)):
if authorization != f"Bearer {API_KEY}":
raise HTTPException(status_code=401, detail="Unauthorized")
return authorization
class Review(BaseModel):
website: str
rating: int
text: str
date: str
@app.get("/review")
def get_reviews(auth=Depends(get_api_key)):
conn = sqlite3.connect("database/reviews.db")
cur = conn.cursor()
cur.execute("""
SELECT * FROM reviews;
""")
data = cur.fetchall()
print(data)
conn.commit()
conn.close()
return {"data": data}
@app.post("/review")
def post_review(review: Review):
conn = sqlite3.connect("database/reviews.db")
cur = conn.cursor()
cur.execute("""
INSERT INTO reviews (website, rating, text, created_at)
VALUES (?, ?, ?, ?)
""", (review.website, review.rating, review.text, review.date))
conn.commit()
conn.close()
return {"status": "ok"}
@app.put("/review")
def put_review(auth=Depends(get_api_key)):
conn = sqlite3.connect("database/reviews.db")
cur = conn.cursor()
cur.execute("""
CREATE TABLE IF NOT EXISTS reviews (
id INTEGER PRIMARY KEY AUTOINCREMENT,
website TEXT,
rating INTEGER,
text TEXT,
created_at TEXT
)
""")
conn.commit()
conn.close()
return {"status": "ok"}
@app.delete("/review")
def post_review(auth=Depends(get_api_key)):
conn = sqlite3.connect("database/reviews.db")
cur = conn.cursor()
cur.execute("""
DELETE FROM reviews;
""")
conn.commit()
conn.close()
return {"status": "ok"}