codeby-hp/human-pose-classification
1
1import logging2import os3import time4import warnings5from pathlib import Path6 7import torch8from fastapi import FastAPI, UploadFile, File, HTTPException9from fastapi.responses import HTMLResponse10from fastapi.templating import Jinja2Templates11from fastapi.requests import Request12from transformers import AutoImageProcessor, pipeline13from PIL import Image14import io15 16from scripts.data_model import (17 PoseClassificationResponse,18 PosePrediction,19)20from scripts.huggingface_load import download_model_from_huggingface21 22USE_HUGGINGFACE_MODELS = True23 24warnings.filterwarnings("ignore")25 26# Configure logging27logging.basicConfig(level=logging.INFO)28logger = logging.getLogger(__name__)29 30# Initialize FastAPI app31app = FastAPI(32 title="Pose Classification API",33 description="ViT-based human pose classification service",34 version="0.0.0",35)36 37# Setup templates38template_dir = Path(__file__).parent / "templates"39if template_dir.exists():40 templates = Jinja2Templates(directory=str(template_dir))41 42# Device selection43device = torch.device("cuda" if torch.cuda.is_available() else "cpu")44logger.info(f"Using device: {device}")45 46# Model initialization47MODEL_NAME = "vit-human-pose-classification"48LOCAL_MODEL_PATH = f"ml-models/{MODEL_NAME}"49FORCE_DOWNLOAD = False50 51# Global model variables52pose_model = None53image_processor = None54 55 56def initialize_model():57 """Initialize the pose classification model."""58 global pose_model, image_processor59 60 try:61 logger.info("Initializing pose classification model...")62 63 # Download model if not present64 if not os.path.isdir(LOCAL_MODEL_PATH) or FORCE_DOWNLOAD:65 if USE_HUGGINGFACE_MODELS:66 logger.info(f"Downloading model from Hugging Face to {LOCAL_MODEL_PATH}")67 success = download_model_from_huggingface(LOCAL_MODEL_PATH)68 else:69 logger.info("failed to download model")70 71 if not success:72 logger.error("Failed to download model")73 return False74 75 # Load image processor76 image_processor = AutoImageProcessor.from_pretrained(77 LOCAL_MODEL_PATH,78 use_fast=True,79 local_files_only=True,80 )81 82 # Load model pipeline83 pose_model = pipeline(84 "image-classification",85 model=LOCAL_MODEL_PATH,86 device=device,87 image_processor=image_processor,88 )89 90 logger.info("Model initialized successfully")91 return True92 93 except Exception as e:94 logger.error(f"Error initializing model: {e}")95 return False96 97 98@app.on_event("startup")99async def startup_event():100 """Initialize model on startup."""101 if not initialize_model():102 logger.warning("Model initialization failed, app will not be fully functional")103 104 105@app.get("/", response_class=HTMLResponse)106async def read_root(request: Request):107 """Serve the main UI page."""108 if template_dir.exists():109 return templates.TemplateResponse("index.html", {"request": request})110 return """111 <!DOCTYPE html>112 <html>113 <head><title>Pose Classification</title></head>114 <body><p>Template not found</p></body>115 </html>116 """117 118 119@app.get("/health")120async def health_check():121 """Health check endpoint."""122 if pose_model is not None:123 return {"status": "healthy", "model_loaded": True}124 return {"status": "unhealthy", "model_loaded": False}125 126 127@app.post("/api/v1/classify")128async def classify_pose(file: UploadFile = File(...)) -> PoseClassificationResponse:129 """Classify pose from uploaded image.130 131 Args:132 file: Image file to classify133 134 Returns:135 PoseClassificationResponse with prediction results136 """137 if pose_model is None:138 raise HTTPException(139 status_code=503,140 detail="Model not loaded. Please try again later.",141 )142 143 try:144 # Read and validate image145 content = await file.read()146 image = Image.open(io.BytesIO(content))147 148 # Run inference149 start_time = time.time()150 output = pose_model(image)151 inference_time = int((time.time() - start_time) * 1000)152 153 # Extract top prediction154 top_prediction = output[0]155 156 return PoseClassificationResponse(157 prediction=PosePrediction(158 label=top_prediction["label"],159 score=round(top_prediction["score"], 4),160 ),161 prediction_time_ms=inference_time,162 )163 164 except Exception as e:165 logger.error(f"Error during inference: {e}")166 raise HTTPException(167 status_code=400,168 detail=f"Error processing image: {str(e)}",169 )170 171if __name__=="__main__":172 import uvicorn173 uvicorn.run(app="app:app", port=8000, reload=True, host="0.0.0.0")174 