rajkamerkar's picture
Update app.py
d458590 verified
Raw
History Blame Contribute Delete
6.05 kB
import cv2
import mediapipe as mp
import numpy as np
from mediapipe.tasks import python
from mediapipe.tasks.python import vision
import gradio as gr
# Initialize MediaPipe hair segmentation model
MODEL_PATH = "hair_segmenter.tflite"
base_options = python.BaseOptions(model_asset_path=MODEL_PATH)
options = vision.ImageSegmenterOptions(
base_options=base_options,
running_mode=vision.RunningMode.IMAGE
)
segmenter = vision.ImageSegmenter.create_from_options(options)
def apply_hair_color(frame, mask, color, alpha=0.5, preserve_highlights=0.3):
"""Apply realistic hair color with preserved highlights"""
lab = cv2.cvtColor(frame, cv2.COLOR_BGR2LAB)
l, a, b = cv2.split(lab)
highlight_mask = (l > 180).astype(np.float32) * preserve_highlights
hair_mask = mask.astype(np.uint8)
colored_hair = np.zeros_like(frame)
colored_hair[:] = color[::-1] # BGR format
colored_lab = cv2.cvtColor(colored_hair, cv2.COLOR_BGR2LAB)
colored_lab[:,:,0] = l * (1 + highlight_mask)
blended_lab = cv2.addWeighted(lab, 1-alpha, colored_lab, alpha, 0)
colored_result = cv2.cvtColor(blended_lab, cv2.COLOR_LAB2BGR)
inverse_mask = cv2.bitwise_not(hair_mask)
background = cv2.bitwise_and(frame, frame, mask=inverse_mask)
hair_region = cv2.bitwise_and(colored_result, colored_result, mask=hair_mask)
return cv2.add(background, hair_region)
def process_frame(frame, color_name, custom_color, intensity, highlights):
"""Process webcam frame with hair coloring"""
# Professional hair color presets
hair_colors = {
"Chocolate Brown": [139, 69, 19],
"Light Golden Brown": [222, 184, 135],
"Soft Black": [51, 51, 51],
"Golden Blonde": [218, 165, 32],
"Dark Auburn": [139, 0, 0],
"Deep Brown": [85, 46, 10],
"Light Auburn": [205, 133, 63],
"Medium Auburn": [165, 42, 42],
"Platinum Blonde": [240, 220, 180],
"Sienna Brown": [160, 82, 45]
}
# Get selected color
if color_name == "Custom":
# Convert hex to RGB
selected_color = [int(custom_color[1:][i:i+2], 16) for i in (0, 2, 4)]
else:
selected_color = hair_colors[color_name]
# Convert frame for MediaPipe
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
mp_image = mp.Image(image_format=mp.ImageFormat.SRGB, data=frame_rgb)
# Perform hair segmentation
segmentation_result = segmenter.segment(mp_image)
if segmentation_result.confidence_masks:
# Process mask
mask = segmentation_result.confidence_masks[0].numpy_view()
mask_resized = cv2.resize(mask, (frame.shape[1], frame.shape[0]))
binary_mask = (mask_resized > 0.3).astype(np.uint8) * 255
binary_mask = cv2.bitwise_not(binary_mask)
# Clean up mask
kernel = np.ones((5,5), np.uint8)
binary_mask = cv2.erode(binary_mask, kernel, iterations=1)
binary_mask = cv2.dilate(binary_mask, kernel, iterations=2)
# Apply hair color
colored_frame = apply_hair_color(
frame,
binary_mask,
selected_color,
alpha=intensity,
preserve_highlights=highlights
)
return cv2.cvtColor(colored_frame, cv2.COLOR_BGR2RGB)
return frame_rgb
def create_interface():
"""Create Gradio interface"""
css = """
.my-group {max-width: 800px !important; max-height: 800px !important;}
.my-column {display: flex !important; justify-content: center !important; align-items: center !important;}
"""
with gr.Blocks(css=css) as demo:
gr.Markdown("# Virtual Hair Color Try-On")
with gr.Row():
with gr.Column():
# Color selection controls
color_dropdown = gr.Dropdown(
choices=["Chocolate Brown", "Light Golden Brown", "Soft Black",
"Golden Blonde", "Dark Auburn", "Deep Brown", "Light Auburn",
"Medium Auburn", "Platinum Blonde", "Sienna Brown", "Custom"],
value="Chocolate Brown",
label="Choose Hair Color"
)
custom_color = gr.ColorPicker(
label="Custom Color",
visible=False
)
# Advanced controls
intensity = gr.Slider(
minimum=0.0,
maximum=1.0,
value=0.8,
step=0.1,
label="Color Intensity"
)
highlights = gr.Slider(
minimum=0.0,
maximum=1.0,
value=0.3,
step=0.1,
label="Preserve Natural Highlights"
)
with gr.Column():
# Modified WebRTC component for live preview
webrtc = gr.Image(
label="Live Preview",
source="webcam",
streaming=True,
type="numpy"
)
# Show/hide custom color picker based on selection
def update_color_picker(choice):
return gr.ColorPicker(visible=(choice == "Custom"))
color_dropdown.change(
fn=update_color_picker,
inputs=[color_dropdown],
outputs=[custom_color]
)
# Stream processing with continuous updates
webrtc.stream(
fn=process_frame,
inputs=[
webrtc,
color_dropdown,
custom_color,
intensity,
highlights
],
outputs=[webrtc],
stream_every=0.1 # Process frame every 0.1 seconds for smooth updates
)
return demo
if __name__ == "__main__":
demo = create_interface()
demo.queue()
demo.launch()