513 lines
20 KiB
Python
513 lines
20 KiB
Python
from fastapi import FastAPI, UploadFile, File, Response, Form, HTTPException, BackgroundTasks
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
from fastapi.responses import FileResponse, JSONResponse
|
|
from PIL import Image
|
|
import io
|
|
import shutil
|
|
from rembg import remove as rembg_remove, new_session
|
|
import time
|
|
import numpy as np
|
|
import tempfile
|
|
import uuid
|
|
import os
|
|
import subprocess
|
|
from transformers import pipeline
|
|
from transparent_background import Remover
|
|
import logging
|
|
import asyncio
|
|
from datetime import datetime, timedelta
|
|
import torch
|
|
from .ormbg import ORMBGProcessor
|
|
from typing import Dict
|
|
from contextlib import contextmanager
|
|
|
|
from carvekit.ml.files.models_loc import download_all
|
|
|
|
download_all()
|
|
|
|
from carvekit.ml.wrap.u2net import U2NET
|
|
from carvekit.ml.wrap.basnet import BASNET
|
|
from carvekit.ml.wrap.fba_matting import FBAMatting
|
|
from carvekit.ml.wrap.deeplab_v3 import DeepLabV3
|
|
from carvekit.ml.wrap.tracer_b7 import TracerUniversalB7
|
|
from carvekit.api.interface import Interface
|
|
from carvekit.pipelines.postprocessing import MattingMethod
|
|
from carvekit.pipelines.preprocessing import PreprocessingStub
|
|
from carvekit.trimap.generator import TrimapGenerator
|
|
|
|
|
|
# Set up logging
|
|
logging.basicConfig(level=logging.DEBUG)
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# Add ORMBG model initialization
|
|
ormbg_model_path = os.path.expanduser("~/.ormbg/ormbg.pth")
|
|
try:
|
|
ormbg_processor = ORMBGProcessor(ormbg_model_path)
|
|
if torch.cuda.is_available():
|
|
ormbg_processor.to("cuda")
|
|
else:
|
|
ormbg_processor.to("cpu")
|
|
except FileNotFoundError:
|
|
logger.error(f"ORMBG model file not found: {ormbg_model_path}")
|
|
print("Error: ORMBG model file not found. Please run 'npm run setup-server' to download it.")
|
|
exit(1)
|
|
|
|
app = FastAPI()
|
|
|
|
# Create temp_videos folder if it doesn't exist
|
|
TEMP_VIDEOS_DIR = "temp_videos"
|
|
os.makedirs(TEMP_VIDEOS_DIR, exist_ok=True)
|
|
|
|
# Create a frames directory within temp_videos
|
|
FRAMES_DIR = os.path.join(TEMP_VIDEOS_DIR, "frames")
|
|
os.makedirs(FRAMES_DIR, exist_ok=True)
|
|
|
|
# Add a dictionary to store processing status
|
|
processing_status = {}
|
|
|
|
# Add CORS middleware
|
|
app.add_middleware(
|
|
CORSMiddleware,
|
|
allow_origins=["*"],
|
|
allow_credentials=True,
|
|
allow_methods=["*"],
|
|
allow_headers=["*"],
|
|
)
|
|
|
|
async def cleanup_old_videos():
|
|
while True:
|
|
current_time = datetime.now()
|
|
for item in os.listdir(TEMP_VIDEOS_DIR):
|
|
item_path = os.path.join(TEMP_VIDEOS_DIR, item)
|
|
item_modified = datetime.fromtimestamp(os.path.getmtime(item_path))
|
|
if current_time - item_modified > timedelta(minutes=10):
|
|
if os.path.isfile(item_path):
|
|
os.remove(item_path)
|
|
logger.info(f"Removed old file: {item_path}")
|
|
elif os.path.isdir(item_path):
|
|
shutil.rmtree(item_path)
|
|
logger.info(f"Removed old directory: {item_path}")
|
|
await asyncio.sleep(600) # Run every 10 minutes
|
|
|
|
# Pre-load all models
|
|
bria_model = pipeline("image-segmentation", model="briaai/RMBG-1.4", trust_remote_code=True, device="cpu")
|
|
inspyrenet_model = Remover()
|
|
inspyrenet_model.model.cpu()
|
|
rembg_models = {
|
|
'u2net': new_session('u2net'),
|
|
'u2net_human_seg': new_session('u2net_human_seg'),
|
|
'isnet-general-use': new_session('isnet-general-use'),
|
|
'isnet-anime': new_session('isnet-anime')
|
|
}
|
|
|
|
# Initialize Carvekit models
|
|
def initialize_carvekit_model(seg_pipe_class, device='cuda'):
|
|
model = Interface(
|
|
pre_pipe=PreprocessingStub(),
|
|
post_pipe=MattingMethod(
|
|
matting_module=FBAMatting(device=device, input_tensor_size=2048, batch_size=1),
|
|
trimap_generator=TrimapGenerator(),
|
|
device=device
|
|
),
|
|
seg_pipe=seg_pipe_class(device=device, batch_size=1)
|
|
)
|
|
model.segmentation_pipeline.to('cpu')
|
|
return model
|
|
|
|
carvekit_models = {
|
|
'u2net': initialize_carvekit_model(U2NET),
|
|
'tracer': initialize_carvekit_model(TracerUniversalB7),
|
|
'basnet': initialize_carvekit_model(BASNET),
|
|
'deeplab': initialize_carvekit_model(DeepLabV3)
|
|
}
|
|
|
|
# ---- Method registry + aliases (NEW) -----------------------------------------
|
|
CARVEKIT = list(carvekit_models.keys()) # ['u2net', 'tracer', 'basnet', 'deeplab']
|
|
REMBG = ['u2net_human_seg', 'isnet-general-use', 'isnet-anime']
|
|
EXTRA = ['ormbg', 'bria', 'inspyrenet']
|
|
ALLOWED_METHODS = sorted(set(CARVEKIT + REMBG + EXTRA))
|
|
|
|
ALIASES = {
|
|
'deeplabv3': 'deeplab',
|
|
'tracer-b7': 'tracer',
|
|
'tracer_b7': 'tracer',
|
|
'u2net_human': 'u2net_human_seg',
|
|
'isnet_general_use': 'isnet-general-use',
|
|
'bria14': 'bria',
|
|
# (intentionally no alias for 'rembg'/'open-rmbg')
|
|
}
|
|
|
|
@app.get("/methods")
|
|
def list_methods():
|
|
return {"methods": ALLOWED_METHODS}
|
|
|
|
@app.get("/health")
|
|
def health():
|
|
return JSONResponse({"ok": True})
|
|
# -----------------------------------------------------------------------------
|
|
|
|
# Ensure GPU memory is cleared after initialization
|
|
torch.cuda.empty_cache()
|
|
|
|
def process_with_bria(image):
|
|
result = bria_model(image, return_mask=True)
|
|
mask = result
|
|
if not isinstance(mask, Image.Image):
|
|
mask = Image.fromarray((mask * 255).astype('uint8'))
|
|
no_bg_image = Image.new("RGBA", image.size, (0, 0, 0, 0))
|
|
no_bg_image.paste(image, mask=mask)
|
|
return no_bg_image
|
|
|
|
def process_with_ormbg(image):
|
|
result = ormbg_processor.process_image(image)
|
|
return result
|
|
|
|
def process_with_inspyrenet(image):
|
|
return inspyrenet_model.process(image, type='rgba')
|
|
|
|
def process_with_rembg(image, model='u2net'):
|
|
return rembg_remove(image, session=rembg_models[model])
|
|
|
|
def process_with_carvekit(image, model='u2net'):
|
|
# Initialize segmentation network based on model input
|
|
if model == 'u2net':
|
|
seg_net = U2NET(device='cuda', batch_size=1)
|
|
elif model == 'tracer':
|
|
seg_net = TracerUniversalB7(device='cuda', batch_size=1)
|
|
elif model == 'basnet':
|
|
seg_net = BASNET(device='cuda', batch_size=1)
|
|
elif model == 'deeplab':
|
|
seg_net = DeepLabV3(device='cuda', batch_size=1)
|
|
else:
|
|
raise ValueError("Unsupported model type")
|
|
|
|
# Setup the post-processing components
|
|
fba = FBAMatting(device='cuda', input_tensor_size=2048, batch_size=1)
|
|
trimap = TrimapGenerator()
|
|
preprocessing = PreprocessingStub()
|
|
postprocessing = MattingMethod(matting_module=fba, trimap_generator=trimap, device='cuda')
|
|
|
|
interface = Interface(pre_pipe=preprocessing, post_pipe=postprocessing, seg_pipe=seg_net)
|
|
processed_image = interface([image])[0]
|
|
|
|
return processed_image
|
|
|
|
@contextmanager
|
|
def inspyrenet_video_model_context():
|
|
try:
|
|
model = Remover()
|
|
model.model.cuda()
|
|
yield model
|
|
finally:
|
|
model.model.cpu()
|
|
del model
|
|
torch.cuda.empty_cache()
|
|
|
|
@contextmanager
|
|
def carvekit_video_model_context(model_name):
|
|
try:
|
|
if model_name == 'u2net':
|
|
seg_net = U2NET(device='cuda', batch_size=1)
|
|
elif model_name == 'tracer':
|
|
seg_net = TracerUniversalB7(device='cuda', batch_size=1)
|
|
elif model_name == 'basnet':
|
|
seg_net = BASNET(device='cuda', batch_size=1)
|
|
elif model_name == 'deeplab':
|
|
seg_net = DeepLabV3(device='cuda', batch_size=1)
|
|
else:
|
|
raise ValueError("Unsupported model type")
|
|
|
|
fba = FBAMatting(device='cuda', input_tensor_size=2048, batch_size=1)
|
|
trimap = TrimapGenerator()
|
|
preprocessing = PreprocessingStub()
|
|
postprocessing = MattingMethod(matting_module=fba, trimap_generator=trimap, device='cuda')
|
|
|
|
interface = Interface(pre_pipe=preprocessing, post_pipe=postprocessing, seg_pipe=seg_net)
|
|
yield interface
|
|
finally:
|
|
del seg_net, fba, trimap, preprocessing, postprocessing, interface
|
|
torch.cuda.empty_cache()
|
|
|
|
|
|
# Create a global lock for GPU operations
|
|
gpu_lock = asyncio.Lock()
|
|
|
|
@app.post("/remove_background/")
|
|
async def remove_background(file: UploadFile = File(...), method: str = Form(...)):
|
|
try:
|
|
# normalize aliases
|
|
method = ALIASES.get(method, method)
|
|
|
|
# early reject unknown methods with 400, not 500
|
|
if method not in ALLOWED_METHODS:
|
|
raise HTTPException(status_code=400, detail="Invalid method")
|
|
|
|
image_data = await file.read()
|
|
image = Image.open(io.BytesIO(image_data)).convert('RGB')
|
|
|
|
start_time = time.time()
|
|
|
|
async def process_image():
|
|
if method == 'bria':
|
|
return await asyncio.to_thread(process_with_bria, image)
|
|
elif method == 'inspyrenet':
|
|
async with gpu_lock:
|
|
try:
|
|
inspyrenet_model.model.to('cuda')
|
|
result = await asyncio.to_thread(inspyrenet_model.process, image, type='rgba')
|
|
finally:
|
|
inspyrenet_model.model.to('cpu')
|
|
return result
|
|
elif method in ['u2net_human_seg', 'isnet-general-use', 'isnet-anime']:
|
|
return await asyncio.to_thread(process_with_rembg, image, model=method)
|
|
elif method == 'ormbg':
|
|
return await asyncio.to_thread(process_with_ormbg, image)
|
|
elif method in ['u2net', 'tracer', 'basnet', 'deeplab']:
|
|
async with gpu_lock:
|
|
try:
|
|
carvekit_models[method].segmentation_pipeline.to('cuda')
|
|
result = await asyncio.to_thread(carvekit_models[method], [image])
|
|
finally:
|
|
carvekit_models[method].segmentation_pipeline.to('cpu')
|
|
return result[0]
|
|
else:
|
|
raise HTTPException(status_code=400, detail="Invalid method")
|
|
|
|
no_bg_image = await process_image()
|
|
|
|
process_time = time.time() - start_time
|
|
print(f"Background removal time ({method}): {process_time:.2f} seconds")
|
|
|
|
async with gpu_lock:
|
|
torch.cuda.empty_cache()
|
|
|
|
with io.BytesIO() as output:
|
|
no_bg_image.save(output, format="PNG")
|
|
content = output.getvalue()
|
|
|
|
return Response(content=content, media_type="image/png")
|
|
|
|
except HTTPException as he:
|
|
# let 4xx bubble as-is
|
|
raise he
|
|
except Exception as e:
|
|
print(str(e))
|
|
raise HTTPException(status_code=500, detail=str(e))
|
|
|
|
async def process_frame(frame_path, method):
|
|
img = Image.open(frame_path).convert('RGB')
|
|
|
|
if method == 'bria':
|
|
processed_frame = await asyncio.to_thread(process_with_bria, img)
|
|
elif method in ['u2net_human_seg', 'isnet-general-use', 'isnet-anime']:
|
|
processed_frame = await asyncio.to_thread(process_with_rembg, img, model=method)
|
|
elif method == 'ormbg':
|
|
processed_frame = await asyncio.to_thread(process_with_ormbg, img)
|
|
else:
|
|
raise ValueError("Invalid method")
|
|
|
|
return processed_frame
|
|
|
|
async def process_video(video_path, method, video_id):
|
|
try:
|
|
processing_status[video_id] = {'status': 'processing', 'progress': 0, 'message': 'Initializing'}
|
|
|
|
logger.info(f"Starting video processing: {video_path}")
|
|
logger.info(f"Method: {method}")
|
|
logger.info(f"Video ID: {video_id}")
|
|
|
|
|
|
# Check video frame count
|
|
frame_count_command = ['ffmpeg.ffprobe', '-v', 'error', '-select_streams', 'v:0', '-count_packets',
|
|
'-show_entries', 'stream=nb_read_packets', '-of', 'csv=p=0', video_path]
|
|
process = await asyncio.create_subprocess_exec(
|
|
*frame_count_command,
|
|
stdout=asyncio.subprocess.PIPE,
|
|
stderr=asyncio.subprocess.PIPE
|
|
)
|
|
stdout, stderr = await process.communicate()
|
|
|
|
if process.returncode != 0:
|
|
logger.error(f"Error counting frames: {stderr.decode()}")
|
|
processing_status[video_id] = {'status': 'error', 'message': 'Error counting frames'}
|
|
return
|
|
|
|
frame_count = int(stdout.decode().strip())
|
|
logger.info(f"Video frame count: {frame_count}")
|
|
|
|
#DISABLED VIDEO LENGTH LIMIT
|
|
#if frame_count > 250:
|
|
# logger.warning(f"Video too long: {frame_count} frames")
|
|
# processing_status[video_id] = {'status': 'error', 'message': 'Video too long (max 250 frames)'}
|
|
# return
|
|
|
|
# Create a unique directory for this video's frames
|
|
frames_dir = os.path.join(FRAMES_DIR, video_id)
|
|
os.makedirs(frames_dir, exist_ok=True)
|
|
logger.info(f"Created frames directory: {frames_dir}")
|
|
|
|
# Extract frames from video
|
|
processing_status[video_id] = {'status': 'processing', 'progress': 0, 'message': 'Extracting frames'}
|
|
extract_command = ['ffmpeg', '-i', video_path, f'{frames_dir}/frame_%05d.png']
|
|
logger.info(f"Executing frame extraction command: {' '.join(extract_command)}")
|
|
process = await asyncio.create_subprocess_exec(
|
|
*extract_command,
|
|
stdout=asyncio.subprocess.PIPE,
|
|
stderr=asyncio.subprocess.PIPE
|
|
)
|
|
stdout, stderr = await process.communicate()
|
|
|
|
if process.returncode != 0:
|
|
logger.error(f"Error extracting frames: {stderr.decode()}")
|
|
processing_status[video_id] = {'status': 'error', 'message': 'Error extracting frames'}
|
|
return
|
|
|
|
# Process frames
|
|
processing_status[video_id] = {'status': 'processing', 'progress': 0, 'message': 'Removing background'}
|
|
frame_files = sorted([f for f in os.listdir(frames_dir) if f.endswith('.png')])
|
|
total_frames = len(frame_files)
|
|
logger.info(f"Number of extracted frames: {total_frames}")
|
|
|
|
if total_frames == 0:
|
|
logger.error("No frames were extracted from the video")
|
|
processing_status[video_id] = {'status': 'error', 'message': 'No frames were extracted from the video'}
|
|
return
|
|
|
|
# Initialize the model once, outside the batch processing loop
|
|
if method == 'inspyrenet':
|
|
print("start init")
|
|
model_context = inspyrenet_video_model_context()
|
|
print("start enter")
|
|
model = model_context.__enter__()
|
|
print("finish enter")
|
|
elif method in ['u2net', 'tracer', 'basnet', 'deeplab']:
|
|
model_context = carvekit_video_model_context(method)
|
|
model = model_context.__enter__()
|
|
else:
|
|
model = None # For other methods that don't require a specific model
|
|
|
|
try:
|
|
async def process_frame_batch(start_idx, end_idx):
|
|
for i in range(start_idx, min(end_idx, total_frames)):
|
|
frame_file = frame_files[i]
|
|
frame_path = os.path.join(frames_dir, frame_file)
|
|
img = Image.open(frame_path).convert('RGB')
|
|
|
|
if method == 'inspyrenet':
|
|
processed_frame = model.process(img, type='rgba')
|
|
elif method in ['u2net', 'tracer', 'basnet', 'deeplab']:
|
|
processed_frame = model([img])[0]
|
|
else:
|
|
processed_frame = await process_frame(frame_path, method)
|
|
|
|
processed_frame.save(frame_path, format='PNG')
|
|
progress = (i + 1) / total_frames * 100
|
|
processing_status[video_id] = {'status': 'processing', 'progress': progress}
|
|
|
|
batch_size = 3
|
|
for i in range(0, total_frames, batch_size):
|
|
await process_frame_batch(i, i + batch_size)
|
|
await asyncio.sleep(0) # Allow other tasks to run
|
|
|
|
finally:
|
|
# Ensure we clean up the model context
|
|
if method in ['inspyrenet', 'u2net', 'tracer', 'basnet', 'deeplab']:
|
|
model_context.__exit__(None, None, None)
|
|
|
|
# Create output video
|
|
processing_status[video_id] = {'status': 'processing', 'progress': 100, 'message': 'Encoding video'}
|
|
output_path = os.path.join(TEMP_VIDEOS_DIR, f"output_{video_id}.webm")
|
|
create_video_command = [
|
|
'ffmpeg',
|
|
'-framerate', '24',
|
|
'-i', f'{frames_dir}/frame_%05d.png',
|
|
'-c:v', 'libvpx-vp9',
|
|
'-pix_fmt', 'yuva420p',
|
|
'-lossless', '1',
|
|
output_path
|
|
]
|
|
logger.info(f"Executing video creation command: {' '.join(create_video_command)}")
|
|
process = await asyncio.create_subprocess_exec(
|
|
*create_video_command,
|
|
stdout=asyncio.subprocess.PIPE,
|
|
stderr=asyncio.subprocess.PIPE
|
|
)
|
|
stdout, stderr = await process.communicate()
|
|
|
|
if process.returncode != 0:
|
|
logger.error(f"Error creating output video: {stderr.decode()}")
|
|
processing_status[video_id] = {'status': 'error', 'message': 'Error creating output video'}
|
|
return
|
|
|
|
logger.info(f"Video processing completed. Output path: {output_path}")
|
|
processing_status[video_id] = {'status': 'completed', 'output_path': output_path}
|
|
|
|
except Exception as e:
|
|
logger.exception("Error in video processing")
|
|
processing_status[video_id] = {'status': 'error', 'message': str(e)}
|
|
finally:
|
|
torch.cuda.empty_cache()
|
|
|
|
# Clean up frames directory
|
|
for file in os.listdir(frames_dir):
|
|
os.remove(os.path.join(frames_dir, file))
|
|
os.rmdir(frames_dir)
|
|
logger.info(f"Cleaned up frames directory: {frames_dir}")
|
|
|
|
@app.post("/remove_background_video/")
|
|
async def remove_background_video(background_tasks: BackgroundTasks, file: UploadFile = File(...), method: str = Form(...)):
|
|
try:
|
|
logger.info(f"Starting video background removal with method: {method}")
|
|
|
|
# Generate a unique filename for the uploaded video
|
|
video_id = str(uuid.uuid4())
|
|
filename = f"input_{video_id}.mp4"
|
|
file_path = os.path.join(TEMP_VIDEOS_DIR, filename)
|
|
|
|
# Save uploaded video to the temp_videos folder
|
|
with open(file_path, "wb") as buffer:
|
|
content = await file.read()
|
|
buffer.write(content)
|
|
|
|
logger.info(f"Video file saved: {file_path}")
|
|
logger.info(f"File exists: {os.path.exists(file_path)}")
|
|
logger.info(f"File size: {os.path.getsize(file_path)} bytes")
|
|
|
|
if not os.path.exists(file_path):
|
|
raise HTTPException(status_code=500, detail=f"Failed to create video file: {file_path}")
|
|
|
|
# Start processing in the background
|
|
background_tasks.add_task(process_video, file_path, method, video_id)
|
|
|
|
return {"video_id": video_id}
|
|
|
|
except Exception as e:
|
|
logger.exception(f"Error in video processing: {str(e)}")
|
|
raise HTTPException(status_code=500, detail=f"Error in video processing: {str(e)}")
|
|
|
|
@app.get("/status/{video_id}")
|
|
async def get_status(video_id: str):
|
|
if video_id not in processing_status:
|
|
raise HTTPException(status_code=404, detail="Video ID not found")
|
|
|
|
status = processing_status[video_id]
|
|
|
|
if status['status'] == 'completed':
|
|
output_path = status['output_path']
|
|
if not os.path.exists(output_path):
|
|
raise HTTPException(status_code=404, detail="Processed video file not found")
|
|
|
|
return FileResponse(output_path, media_type="video/webm", filename=f"processed_video_{video_id}.webm")
|
|
|
|
return status
|
|
|
|
@app.on_event("startup")
|
|
async def startup_event():
|
|
asyncio.create_task(cleanup_old_videos())
|
|
|
|
|
|
if __name__ == "__main__":
|
|
import uvicorn
|
|
uvicorn.run(app, host="0.0.0.0", port=9876)
|
|
|