AI-Face-Swap / app.py
chunpeng xiang
Add ZhiQITeAI hyperlink
a81e14c
Raw History Blame Contribute Delete
12.8 kB
import os
import io
import time
import uuid
import base64
import logging
import asyncio
import traceback
from typing import Optional, Dict, Any, Union
import requests
from PIL import Image
import gradio as gr
import aiohttp
from aiohttp import ClientTimeout
# Configure logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)
if not os.environ.get("SPACE_ID"):
try:
from dotenv import load_dotenv
load_dotenv()
except Exception as e:
logger.error(f"Failed to load environment variables: {str(e)}")
# API configuration
API_CONFIG = {
# Get API endpoints from environment variables, use default values if not exist
"SUBMIT_API_URL": os.environ.get("FACESWAP_SUBMIT_API", ""),
"RESULT_API_URL": os.environ.get("FACESWAP_RESULT_API", ""),
"SUBMIT_API_KEY_ID": os.environ.get("FACESWAP_API_KEY_ID", ""),
"SUBMIT_API_KEY_SECRET": os.environ.get("FACESWAP_API_KEY_SECRET", ""),
"RESULT_API_KEY_ID": os.environ.get("RESULT_API_KEY_ID", ""),
"RESULT_API_KEY_SECRET": os.environ.get("RESULT_API_KEY_SECRET", ""),
"TIMEOUT": int(os.environ.get("FACESWAP_TIMEOUT", "26")), # Polling timeout (seconds)
"POLL_INTERVAL": int(os.environ.get("FACESWAP_POLL_INTERVAL", "2")), # Polling interval (seconds)
}
def image_to_base64(img: Union[str, Image.Image]) -> str:
"""
Convert an image to a base64 encoded string with MIME type prefix
Parameters:
img: Image file path or PIL.Image object
Returns:
Base64 encoded string with MIME type prefix
"""
if isinstance(img, str):
# If input is a file path
img_extension = os.path.splitext(img)[-1].lower()
mime_type_map = {
'.jpg': 'image/jpeg',
'.jpeg': 'image/jpeg',
'.png': 'image/png',
'.gif': 'image/gif',
'.bmp': 'image/bmp',
'.webp': 'image/webp'
}
mime_type = mime_type_map.get(img_extension, 'image/jpeg') # Default to jpeg
# Open image file and convert to base64
with open(img, 'rb') as img_file:
img_data = base64.b64encode(img_file.read()).decode('utf-8')
else:
# If input is a PIL.Image object
buffered = io.BytesIO()
img_format = img.format if img.format else 'PNG'
mime_type = f'image/{img_format.lower()}' if img_format else 'image/png'
img.save(buffered, format=img_format if img_format else 'PNG')
img_data = base64.b64encode(buffered.getvalue()).decode('utf-8')
# Add MIME type prefix
return f"data:{mime_type};base64,{img_data}"
def submit_face_swap_task(source_base64: str, target_base64: str, input_task_id: str) -> Optional[str]:
"""Submit face swap task to async API, return task ID"""
try:
api_url = API_CONFIG["SUBMIT_API_URL"]
payload = {
"sourceUrl": source_base64,
"targetUrl": target_base64,
"businessTaskId": input_task_id,
"imageDetect": True,
"imageBase64Format": True
}
# Add API keys to request headers if available
headers = {
'X-Request-Req-Accesskeyid': API_CONFIG["SUBMIT_API_KEY_ID"],
'X-Request-Req-Accesskeysecret': API_CONFIG["SUBMIT_API_KEY_SECRET"],
'Content-Type': 'application/json'
}
response = requests.post(api_url, json=payload, headers=headers, timeout=10)
response.raise_for_status()
result = response.json()
# Extract task ID based on actual API response format
task_id = result.get("data", {}).get("task_id")
logger.info(f"Successfully submitted face swap task, task ID: {task_id}")
return task_id
except Exception as e:
logger.exception(f"Failed to submit face swap task: {str(e)}")
return None
def poll_face_swap_result(task_id: str, timeout: int = 30) -> Optional[Dict[str, Any]]:
"""Poll for face swap task result until success or timeout"""
start_time = time.time()
poll_interval = 2 # Polling interval in seconds
api_url = API_CONFIG["RESULT_API_URL"]
payload = {
"saasTaskId": task_id
}
# Add API keys to request headers if available
headers = {
'X-Request-Req-Accesskeyid': API_CONFIG["RESULT_API_KEY_ID"],
'X-Request-Req-Accesskeysecret': API_CONFIG["RESULT_API_KEY_SECRET"],
'Content-Type': 'application/json'
}
while time.time() - start_time < timeout:
try:
response = requests.post(api_url, json=payload, headers=headers, timeout=2)
response.raise_for_status()
result = response.json()
# Check if task is completed
if result:
if result.get("code", None) == "200":
data = result.get("data", None)
if data is None:
time.sleep(poll_interval)
else:
return data
else:
logger.error(f"Task failed")
return None
else:
# Task not completed, wait and continue polling
time.sleep(poll_interval)
except Exception as e:
logger.warning(f"Error during polling: {str(e)}")
time.sleep(poll_interval)
logger.error(f"Polling timeout, task ID: {task_id}")
return None
def download_image(url: str) -> Optional[Image.Image]:
"""Download image from URL"""
try:
response = requests.get(url, timeout=10)
response.raise_for_status()
image_data = io.BytesIO(response.content)
image = Image.open(image_data)
logger.info(f"Successfully downloaded result image, URL: {url}")
return image
except Exception as e:
logger.exception(f"Failed to download image: {str(e)}")
return None
async def download_image_async(encoding: str, logger: logging.Logger, max_retries: int=2,
total: int=7, connect: float=1.5, sock_read: float=5) -> Image.Image:
timeout = ClientTimeout(total=total, connect=connect, sock_read=sock_read)
async with aiohttp.ClientSession(timeout=timeout) as session:
retries = 0
while retries < max_retries:
try:
if encoding.startswith("http://") or encoding.startswith("https://"):
t = time.time()
async with session.get(encoding) as response:
response.raise_for_status() # Check HTTP status code
content = await response.read() # Asynchronously read response content
image = Image.open((io.BytesIO(content)))
logger.info(f"Image with link {encoding} downloaded, costs time {time.time() - t}s")
return image
else:
logger.error(f"Download failed: {encoding}, it is not a downloadable url")
return None
except TimeoutError as e:
retries += 1
logger.warning(f"Attempt {retries} failed with error: {e}")
if retries >= max_retries:
logger.error(f"Download {encoding} failed after {max_retries} attempts.")
return None
except Exception as e:
logger.error(f'download {encoding} failed with {e}, {traceback.format_exc()}')
return None
# Face swap processing function - calling async API
def face_swap(source_img: str, target_img: str) -> Optional[str]:
"""
Synthesize facial features from source image onto target image using async API
Parameters:
source_img: Source image (containing the face to extract)
target_img: Target image (image whose face will be replaced)
Returns:
Face-swapped image, or None if failed
"""
# Input validation
if source_img is None or target_img is None:
logger.warning("Source image or target image is empty")
return None, "Source image or target image is empty"
try:
# Generate task ID
input_task_id = str(uuid.uuid4())
# Convert images to base64 encoding
source_base64 = image_to_base64(source_img)
target_base64 = image_to_base64(target_img)
# Step 1: Call async API to initiate face swap task
task_id = submit_face_swap_task(source_base64, target_base64, input_task_id)
if task_id is None:
logger.error("Unable to get face swap task ID")
return None, "Unable to get face swap task ID"
time.sleep(4)
# Step 2: Poll for results with 30-second timeout
timeout = API_CONFIG["TIMEOUT"]
output_url = poll_face_swap_result(task_id, timeout=timeout)
if output_url is None:
logger.error("Could not get face swap result within specified time")
return None, "Could not get face swap result within specified time"
# Step 3: Return result image
if output_url == "IMAGE_DETECT_ILLEGAL":
logger.error("Illegal image detected")
image = Image.open("example/NSFW.jpg")
return image, "Illegal image detected"
else:
image = asyncio.run(download_image_async(output_url, logger))
if image is None:
logger.error("Failed to download result image")
return None, "Failed to download result image"
else:
return image, "Face swap successful"
except Exception as e:
logger.exception(f"Error during face swap process: {str(e)}")
return None, "Error during face swap process"
# Create Gradio interface
with gr.Blocks(theme=gr.themes.Default()) as demo:
# Title and introduction
gr.Markdown("# AI Face Swap")
gr.Markdown("The AI Face Swap Developed by [ZhiQITeAI](https://huggingface.co/izhiqiteai) enables seamless and accurate face swapping in images. yangzhi@zhiqiteai.cn for business inquiries.")
with gr.Row():
with gr.Column():
# Source image upload area
source_image = gr.Image(label="Source Image", type="filepath", elem_id="source_image")
# Target image upload area
target_image = gr.Image(label="Target Image", type="filepath", elem_id="target_image")
with gr.Column():
# Output image display area
output_image = gr.Image(label="Output", type="pil")
# Status message
status_message = gr.Markdown("Ready. Please upload source and target images, then click Submit.", visible=True)
# Button area
with gr.Row():
clear_btn = gr.Button("Clear", variant="secondary")
submit_btn = gr.Button("Submit", variant="primary")
with gr.Row():
with gr.Column():
gr.Markdown("## Source Image Examples")
source_exm = gr.Examples(
examples=[
os.path.join("example", "source", file)
for file in os.listdir(os.path.join("example", "source"))
],
examples_per_page=4,
inputs=source_image
)
with gr.Row():
with gr.Column():
gr.Markdown("## Target Image Examples")
target_exm = gr.Examples(
examples=[
os.path.join("example", "target", file)
for file in os.listdir(os.path.join("example", "target"))
],
examples_per_page=4,
inputs=target_image
)
# Set up button events
def clear_inputs():
return None, None, None, "Ready. Please upload source and target images, then click Submit."
clear_btn.click(
fn=clear_inputs,
inputs=None,
outputs=[source_image, target_image, output_image, status_message]
)
submit_btn.click(
fn=face_swap,
inputs=[source_image, target_image],
outputs=[output_image, status_message]
)
# Launch Gradio application
if __name__ == "__main__":
demo.queue(max_size=30)
demo.launch()