Spaces:
Running
Running
Download app.py from zhiqiteai/AI-Face-Swap: direct link, hf CLI and curl.
- Browser
- Download file 12.8 kB
-
https://huggingface.co/spaces/zhiqiteai/AI-Face-Swap/resolve/main/app.py
- Command line
-
hf download hf://spaces/zhiqiteai/AI-Face-Swap/app.py
-
curl -L -o app.py https://huggingface.co/spaces/zhiqiteai/AI-Face-Swap/resolve/main/app.py
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() |