Team Ai
Apppublic

codeby-hp/human-pose-classification

sourceHugging Faceupdated 10mo agoView on Hugging Face
1likes
app.py174 linesDownload Raw Back to fastapi_app
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