added review sql database
This commit is contained in:
+3
-1
@@ -1,4 +1,6 @@
|
|||||||
api/data/
|
api/data/
|
||||||
api/saved_models/
|
api/saved_models/
|
||||||
api/testimages/
|
api/testimages/
|
||||||
api/__pycache__/
|
api/__pycache__/
|
||||||
|
|
||||||
|
**/*.env
|
||||||
+1
-1
@@ -134,7 +134,7 @@ def trainCNN(model, optimizer, loss_module, train_data_loader, validation_data_l
|
|||||||
save_dir = os.path.join(SAVE_PATH, "bird_cnn")
|
save_dir = os.path.join(SAVE_PATH, "bird_cnn")
|
||||||
os.makedirs(save_dir, exist_ok=True)
|
os.makedirs(save_dir, exist_ok=True)
|
||||||
|
|
||||||
save_path = os.path.join(save_dir, f"bird_cnn{epoch+1}")
|
save_path = os.path.join(save_dir, f"bird_cnn")
|
||||||
torch.save(model.state_dict(), save_path)
|
torch.save(model.state_dict(), save_path)
|
||||||
|
|
||||||
print(f"epoch: {epoch+1} | train accuracy: {int(train_acc * 1000) / 10}% | validation accuracy: {int(val_acc * 1000) / 10}%")
|
print(f"epoch: {epoch+1} | train accuracy: {int(train_acc * 1000) / 10}% | validation accuracy: {int(val_acc * 1000) / 10}%")
|
||||||
|
|||||||
@@ -6,4 +6,7 @@ uvicorn==0.46.0
|
|||||||
python-multipart==0.0.27
|
python-multipart==0.0.27
|
||||||
opencv-python==4.13.0.92
|
opencv-python==4.13.0.92
|
||||||
pillow==12.2.0
|
pillow==12.2.0
|
||||||
pycocotools==2.0.11
|
pycocotools==2.0.11
|
||||||
|
|
||||||
|
pydantic==2.13.3
|
||||||
|
pydantic_core==2.46.3
|
||||||
+64
-2
@@ -1,9 +1,12 @@
|
|||||||
from enum import Enum
|
from enum import Enum
|
||||||
import os
|
import os
|
||||||
|
import sqlite3
|
||||||
|
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
from pydantic import BaseModel
|
||||||
import torch
|
import torch
|
||||||
from torchvision import transforms
|
from torchvision import transforms
|
||||||
from fastapi import FastAPI, File, UploadFile
|
from fastapi import Depends, FastAPI, File, HTTPException, Header, UploadFile
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
import io
|
import io
|
||||||
@@ -72,4 +75,63 @@ async def predict(file: UploadFile = File(...)):
|
|||||||
return {
|
return {
|
||||||
"class": bird_species(cls.item()).name,
|
"class": bird_species(cls.item()).name,
|
||||||
"confidence": confidence.item()
|
"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():
|
||||||
|
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.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"}
|
||||||
+20
-9
@@ -5,6 +5,7 @@ from torchvision import datasets, transforms
|
|||||||
from torch.utils.data import DataLoader, random_split
|
from torch.utils.data import DataLoader, random_split
|
||||||
import matplotlib.pyplot as plt
|
import matplotlib.pyplot as plt
|
||||||
|
|
||||||
|
from util import TransformedSubset, visualizeData
|
||||||
from bird_cnn import Bird_CNN, sample, trainCNN
|
from bird_cnn import Bird_CNN, sample, trainCNN
|
||||||
|
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
@@ -12,7 +13,7 @@ from enum import Enum
|
|||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
SAVE_PATH = "./saved_models"
|
SAVE_PATH = "./saved_models"
|
||||||
IMAGE_SIZE = 64
|
IMAGE_SIZE = 128
|
||||||
|
|
||||||
class bird_species(Enum):
|
class bird_species(Enum):
|
||||||
Common_Kingfisher = 0
|
Common_Kingfisher = 0
|
||||||
@@ -23,24 +24,34 @@ class bird_species(Enum):
|
|||||||
Ruddy_Shelduck = 5
|
Ruddy_Shelduck = 5
|
||||||
Sarus_Crane = 6
|
Sarus_Crane = 6
|
||||||
|
|
||||||
transform = transforms.Compose([
|
transform_augemnt = transforms.Compose([
|
||||||
|
transforms.RandomAffine(
|
||||||
|
degrees=35, # no rotation
|
||||||
|
translate=(0.2, 0.2) # shift up to 20% horizontally/vertically
|
||||||
|
),
|
||||||
transforms.RandomHorizontalFlip(p=0.5),
|
transforms.RandomHorizontalFlip(p=0.5),
|
||||||
transforms.RandomRotation(35),
|
|
||||||
transforms.Resize(IMAGE_SIZE),
|
transforms.Resize(IMAGE_SIZE),
|
||||||
transforms.CenterCrop(IMAGE_SIZE),
|
transforms.CenterCrop(IMAGE_SIZE),
|
||||||
transforms.ToTensor()
|
transforms.ToTensor()
|
||||||
])
|
])
|
||||||
|
|
||||||
dataset = datasets.ImageFolder("data/CUB_200_2011/images", transform=transform)
|
transform = transforms.Compose([
|
||||||
|
transforms.Resize(IMAGE_SIZE),
|
||||||
|
transforms.CenterCrop(IMAGE_SIZE),
|
||||||
|
transforms.ToTensor()
|
||||||
|
])
|
||||||
|
|
||||||
|
dataset = datasets.ImageFolder("data/CUB_200_2011/images")
|
||||||
|
|
||||||
train_size = int(0.8 * len(dataset))
|
train_size = int(0.8 * len(dataset))
|
||||||
val_size = len(dataset) - train_size
|
val_size = len(dataset) - train_size
|
||||||
|
|
||||||
train_dataset, val_dataset = random_split(dataset, [train_size, val_size])
|
train_subset, val_subset = random_split(dataset, [train_size, val_size])
|
||||||
|
train_dataset = TransformedSubset(train_subset, transform_augemnt)
|
||||||
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
|
val_dataset = TransformedSubset(val_subset, transform)
|
||||||
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)
|
|
||||||
|
|
||||||
|
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
|
||||||
|
val_loader = DataLoader(val_dataset, batch_size=64, shuffle=False)
|
||||||
|
|
||||||
device = torch.device("cpu") if not torch.cuda.is_available() else torch.device("cuda:0")
|
device = torch.device("cpu") if not torch.cuda.is_available() else torch.device("cuda:0")
|
||||||
print("Using device", device)
|
print("Using device", device)
|
||||||
@@ -50,7 +61,7 @@ model.to(device)
|
|||||||
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)
|
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)
|
||||||
loss_module = nn.CrossEntropyLoss()
|
loss_module = nn.CrossEntropyLoss()
|
||||||
|
|
||||||
trainCNN(model, optimizer, loss_module, train_loader, val_loader, device, 50, SAVE_PATH=SAVE_PATH, save=True)
|
trainCNN(model, optimizer, loss_module, train_loader, val_loader, device, 500, SAVE_PATH=SAVE_PATH, save=True)
|
||||||
exit()
|
exit()
|
||||||
|
|
||||||
image_path = "test1.jpg"
|
image_path = "test1.jpg"
|
||||||
|
|||||||
+30
-15
@@ -1,22 +1,37 @@
|
|||||||
from PIL import Image
|
from PIL import Image
|
||||||
import os
|
import os
|
||||||
|
from matplotlib import pyplot as plt
|
||||||
|
import torch
|
||||||
|
|
||||||
folder = "data/train"
|
class TransformedSubset(torch.utils.data.Dataset):
|
||||||
|
def __init__(self, subset, transform=None):
|
||||||
|
self.subset = subset
|
||||||
|
self.transform = transform
|
||||||
|
|
||||||
widths = []
|
def __getitem__(self, idx):
|
||||||
heights = []
|
x, y = self.subset[idx]
|
||||||
|
|
||||||
for root, _, files in os.walk(folder):
|
if self.transform:
|
||||||
for file in files:
|
x = self.transform(x)
|
||||||
if file.endswith((".jpg", ".png", ".jpeg")):
|
|
||||||
path = os.path.join(root, file)
|
|
||||||
img = Image.open(path)
|
|
||||||
w, h = img.size
|
|
||||||
widths.append(w)
|
|
||||||
heights.append(h)
|
|
||||||
|
|
||||||
avg_w = sum(widths) / len(widths)
|
return x, y
|
||||||
avg_h = sum(heights) / len(heights)
|
|
||||||
|
|
||||||
print("Average width:", avg_w)
|
def __len__(self):
|
||||||
print("Average height:", avg_h)
|
return len(self.subset)
|
||||||
|
|
||||||
|
|
||||||
|
def visualizeData(dataset):
|
||||||
|
images, labels = next(iter(dataset))
|
||||||
|
|
||||||
|
for i in range(4):
|
||||||
|
img = images[i]
|
||||||
|
|
||||||
|
# Convert tensor shape from [C,H,W] -> [H,W,C]
|
||||||
|
img = img.permute(1, 2, 0)
|
||||||
|
|
||||||
|
plt.figure(figsize=(3,3))
|
||||||
|
plt.imshow(img)
|
||||||
|
plt.title(f"Label: {labels[i].item()}")
|
||||||
|
plt.axis("off")
|
||||||
|
|
||||||
|
plt.show()
|
||||||
@@ -29,6 +29,8 @@ services:
|
|||||||
- cnn_network
|
- cnn_network
|
||||||
pull_policy: never
|
pull_policy: never
|
||||||
container_name: cnn_api
|
container_name: cnn_api
|
||||||
|
volumes:
|
||||||
|
- ./database:/api/database
|
||||||
|
|
||||||
|
|
||||||
cnn_website:
|
cnn_website:
|
||||||
|
|||||||
Binary file not shown.
Reference in New Issue
Block a user