CoolFace
Apppublic

ToadPres/3d-scene-generator

sourceHugging Facemitupdated 9mo agoView on Hugging Face
0likes
main.py194 linesDownload Raw Back to root
1"""23D Scene Generation Demo - FastAPI Backend3Combines Google Gemini 2.5 Flash Image with Apple SHARP for text-to-3D generation.4"""5 6import os7import uuid8import shutil9from pathlib import Path10from contextlib import asynccontextmanager11from typing import Optional12 13from fastapi import FastAPI, HTTPException, Header14from fastapi.middleware.cors import CORSMiddleware15from fastapi.staticfiles import StaticFiles16from pydantic import BaseModel17 18from services.gemini_service import generate_concept_image19from services.sharp_service import generate_3d_scene, convert_ply_to_splat20 21 22# Configuration23STATIC_DIR = Path(__file__).parent / "static"24TEMP_DIR = Path(__file__).parent / "temp"25 26# Create directories immediately (needed for StaticFiles mount)27STATIC_DIR.mkdir(exist_ok=True)28TEMP_DIR.mkdir(exist_ok=True)29 30 31@asynccontextmanager32async def lifespan(app: FastAPI):33    """Application lifespan - setup and teardown."""34    # Pre-download SHARP model weights on startup35    print("=" * 60)36    print("๐Ÿš€ Pre-downloading SHARP model weights...")37    print("   This may take 2-5 minutes on first startup.")38    print("=" * 60)39    40    try:41        import subprocess42        from PIL import Image43        44        # Create a dummy image to trigger model download45        dummy_image = TEMP_DIR / "dummy_preload.png"46        dummy_output = TEMP_DIR / "dummy_output"47        dummy_output.mkdir(exist_ok=True)48        49        img = Image.new('RGB', (256, 256), color='gray')50        img.save(str(dummy_image))51        52        result = subprocess.run(53            ["sharp", "-i", str(dummy_image), "-o", str(dummy_output)],54            capture_output=True,55            text=True,56            timeout=60057        )58        print(f"SHARP preload: {result.stdout}")59        60        # Cleanup61        if dummy_image.exists():62            dummy_image.unlink()63        if dummy_output.exists():64            shutil.rmtree(dummy_output, ignore_errors=True)65            66    except Exception as e:67        print(f"โš ๏ธ SHARP preload warning: {e}")68    69    print("=" * 60)70    print("โœ… Backend ready - SHARP model cached!")71    print("=" * 60)72    73    yield  # Application runs74    75    # Shutdown: Cleanup temp files76    if TEMP_DIR.exists():77        shutil.rmtree(TEMP_DIR)78    print("๐Ÿงน Cleanup complete")79 80 81 82app = FastAPI(83    title="3D Scene Generator",84    description="Generate immersive 3D scenes from text prompts",85    version="1.0.0",86    lifespan=lifespan,87)88 89# CORS - Allow frontend access (local and cloud deployments)90app.add_middleware(91    CORSMiddleware,92    allow_origins=[93        "http://localhost:3000", 94        "http://127.0.0.1:3000",95        "http://localhost:3002", 96        "http://127.0.0.1:3002",97        "https://*.vercel.app",  # Vercel deployments98    ],99    allow_origin_regex=r"https://.*\.vercel\.app",  # Dynamic Vercel preview URLs100    allow_credentials=True,101    allow_methods=["*"],102    allow_headers=["*"],103)104 105# Serve static files (generated PLY/splat files)106app.mount("/static", StaticFiles(directory=str(STATIC_DIR)), name="static")107 108 109# Request/Response Models110class GenerateRequest(BaseModel):111    prompt: str112    113class GenerateResponse(BaseModel):114    success: bool115    ply_url: Optional[str] = None116    image_url: Optional[str] = None117    message: Optional[str] = None118    generation_time_ms: Optional[int] = None119 120 121@app.get("/health")122async def health_check():123    """Health check endpoint."""124    return {"status": "healthy", "service": "3d-scene-generator"}125 126 127@app.post("/api/generate", response_model=GenerateResponse)128async def generate_scene(request: GenerateRequest, x_api_key: str = Header(..., alias="X-API-Key")):129    """130    Generate a 3D scene from a text prompt.131    132    Pipeline:133    1. Enhance prompt with depth/3D optimization keywords134    2. Generate 16:9 concept image via Gemini 2.5 Flash Image135    3. Convert to 3D Gaussian Splatting via Apple SHARP136    4. Convert PLY to .splat format for optimized web rendering137    5. Return URLs for frontend to render138    """139    import time140    start_time = time.time()141    142    # Set the API key from the request header143    os.environ["GOOGLE_API_KEY"] = x_api_key144    145    try:146        # Generate unique session ID147        session_id = str(uuid.uuid4())[:8]148        149        # Step 1: Generate concept image150        image_path = TEMP_DIR / f"{session_id}_concept.png"151        await generate_concept_image(request.prompt, str(image_path))152        153        if not image_path.exists():154            raise HTTPException(status_code=500, detail="Image generation failed")155        156        # Step 2: Generate 3D scene via SHARP157        output_dir = STATIC_DIR / session_id158        output_dir.mkdir(exist_ok=True)159        160        ply_path = await generate_3d_scene(str(image_path), str(output_dir))161        162        if not ply_path or not Path(ply_path).exists():163            raise HTTPException(status_code=500, detail="3D reconstruction failed")164        165        # Step 3: Convert PLY to optimized .splat format166        splat_path = str(output_dir / f"{session_id}.splat")167        final_path = await convert_ply_to_splat(ply_path, splat_path)168        169        # Copy image to static for preview170        static_image_path = output_dir / "concept.png"171        shutil.copy(image_path, static_image_path)172        173        # Calculate time174        elapsed_ms = int((time.time() - start_time) * 1000)175        176        # Return URLs (use splat if conversion succeeded, otherwise PLY)177        scene_filename = Path(final_path).name178        return GenerateResponse(179            success=True,180            ply_url=f"/static/{session_id}/{scene_filename}",181            image_url=f"/static/{session_id}/concept.png",182            generation_time_ms=elapsed_ms,183        )184        185    except HTTPException:186        raise187    except Exception as e:188        raise HTTPException(status_code=500, detail=str(e))189 190 191if __name__ == "__main__":192    import uvicorn193    uvicorn.run("main:app", host="0.0.0.0", port=8000, reload=True)194