🎨 Inpaint tool improved - more shapes, eraser, undos (#225)
Co-authored-by: Anthony Wu <pls-file-gh-issue@users.noreply.github.com>
This commit is contained in:
parent
11beed80df
commit
676b30db14
@ -15,7 +15,7 @@ class MaskCreator:
|
|||||||
self.mask_output_path = image_path.with_name(f"{image_path.stem}_mask").with_suffix(".png")
|
self.mask_output_path = image_path.with_name(f"{image_path.stem}_mask").with_suffix(".png")
|
||||||
|
|
||||||
# Create a window and set up display image
|
# Create a window and set up display image
|
||||||
self.window_name = "MFlux Inpaint Mask Creator - Draw with mouse or trackpad (hot keys: (s)ave, (r)eset, (q)uit"
|
self.window_name = "MFlux Inpaint Mask Creator - Tools: (b)rush, (e)rase, (r)ectangle, s(q)uare, (o)val, (t)riangle | (u)ndo, (s)ave, (esc) quit"
|
||||||
self.display_image = self.original_image.copy()
|
self.display_image = self.original_image.copy()
|
||||||
|
|
||||||
# Create a blank mask the same size as the image
|
# Create a blank mask the same size as the image
|
||||||
@ -26,6 +26,14 @@ class MaskCreator:
|
|||||||
self.brush_size = self.BRUSH_SIZES["5"]
|
self.brush_size = self.BRUSH_SIZES["5"]
|
||||||
self.last_point = None
|
self.last_point = None
|
||||||
|
|
||||||
|
# Tool modes
|
||||||
|
self.tool_mode = "brush" # brush, erase, rectangle, square, oval, triangle
|
||||||
|
self.shape_start_point = None
|
||||||
|
|
||||||
|
# Undo history
|
||||||
|
self.mask_history = []
|
||||||
|
self.max_history = 50
|
||||||
|
|
||||||
self.overlay = np.zeros_like(self.original_image)
|
self.overlay = np.zeros_like(self.original_image)
|
||||||
|
|
||||||
# Update display every N drawing events - lower is more responsive
|
# Update display every N drawing events - lower is more responsive
|
||||||
@ -35,12 +43,35 @@ class MaskCreator:
|
|||||||
# Show the initial display
|
# Show the initial display
|
||||||
self.update_display()
|
self.update_display()
|
||||||
|
|
||||||
|
def save_to_history(self):
|
||||||
|
self.mask_history.append(self.mask.copy())
|
||||||
|
if len(self.mask_history) > self.max_history:
|
||||||
|
self.mask_history.pop(0)
|
||||||
|
|
||||||
|
def undo(self):
|
||||||
|
if self.mask_history:
|
||||||
|
self.mask = self.mask_history.pop()
|
||||||
|
self.update_display()
|
||||||
|
print("Undo completed")
|
||||||
|
else:
|
||||||
|
print("Nothing to undo")
|
||||||
|
|
||||||
def mouse_callback(self, event, x, y, flags, param):
|
def mouse_callback(self, event, x, y, flags, param):
|
||||||
|
if self.tool_mode in ["brush", "erase"]:
|
||||||
|
self._handle_brush(event, x, y)
|
||||||
|
else:
|
||||||
|
self._handle_shape(event, x, y)
|
||||||
|
|
||||||
|
def _handle_brush(self, event, x, y):
|
||||||
|
# Determine color based on tool mode
|
||||||
|
color = 0 if self.tool_mode == "erase" else 255
|
||||||
|
|
||||||
# Start drawing
|
# Start drawing
|
||||||
if event == cv2.EVENT_LBUTTONDOWN:
|
if event == cv2.EVENT_LBUTTONDOWN:
|
||||||
|
self.save_to_history()
|
||||||
self.drawing = True
|
self.drawing = True
|
||||||
self.last_point = (x, y)
|
self.last_point = (x, y)
|
||||||
cv2.circle(self.mask, (x, y), self.brush_size, 255, -1)
|
cv2.circle(self.mask, (x, y), self.brush_size, color, -1)
|
||||||
self.event_counter += 1
|
self.event_counter += 1
|
||||||
if self.event_counter % self.update_frequency == 0:
|
if self.event_counter % self.update_frequency == 0:
|
||||||
self.update_display()
|
self.update_display()
|
||||||
@ -50,9 +81,9 @@ class MaskCreator:
|
|||||||
# Use thickness based on brush size for smoother lines
|
# Use thickness based on brush size for smoother lines
|
||||||
if self.last_point: # Ensure we have a last point
|
if self.last_point: # Ensure we have a last point
|
||||||
# Draw a line between the last point and current point
|
# Draw a line between the last point and current point
|
||||||
cv2.line(self.mask, self.last_point, (x, y), 255, self.brush_size * 2)
|
cv2.line(self.mask, self.last_point, (x, y), color, self.brush_size * 2)
|
||||||
# Also draw a circle at the current point to avoid gaps in fast movements
|
# Also draw a circle at the current point to avoid gaps in fast movements
|
||||||
cv2.circle(self.mask, (x, y), self.brush_size, 255, -1)
|
cv2.circle(self.mask, (x, y), self.brush_size, color, -1)
|
||||||
self.last_point = (x, y)
|
self.last_point = (x, y)
|
||||||
self.event_counter += 1
|
self.event_counter += 1
|
||||||
if self.event_counter % self.update_frequency == 0:
|
if self.event_counter % self.update_frequency == 0:
|
||||||
@ -63,6 +94,83 @@ class MaskCreator:
|
|||||||
self.drawing = False
|
self.drawing = False
|
||||||
self.update_display() # Always update display when stopping
|
self.update_display() # Always update display when stopping
|
||||||
|
|
||||||
|
def _handle_shape(self, event, x, y):
|
||||||
|
if event == cv2.EVENT_LBUTTONDOWN:
|
||||||
|
self.save_to_history()
|
||||||
|
self.shape_start_point = (x, y)
|
||||||
|
self.drawing = True
|
||||||
|
|
||||||
|
elif event == cv2.EVENT_MOUSEMOVE and self.drawing:
|
||||||
|
# Show preview of shape
|
||||||
|
self.update_display()
|
||||||
|
self._draw_shape_preview(self.shape_start_point, (x, y))
|
||||||
|
cv2.imshow(self.window_name, self.display_image)
|
||||||
|
|
||||||
|
elif event == cv2.EVENT_LBUTTONUP and self.drawing:
|
||||||
|
self.drawing = False
|
||||||
|
if self.shape_start_point:
|
||||||
|
self._draw_shape_on_mask(self.shape_start_point, (x, y))
|
||||||
|
self.shape_start_point = None
|
||||||
|
self.update_display()
|
||||||
|
|
||||||
|
def _get_shape_params(self, start, end):
|
||||||
|
"""Calculate shape parameters based on tool mode."""
|
||||||
|
if self.tool_mode == "rectangle":
|
||||||
|
return {"type": "rectangle", "start": start, "end": end}
|
||||||
|
elif self.tool_mode == "square":
|
||||||
|
# Force square aspect ratio
|
||||||
|
width = abs(end[0] - start[0])
|
||||||
|
height = abs(end[1] - start[1])
|
||||||
|
size = max(width, height)
|
||||||
|
# Preserve direction of drag
|
||||||
|
end_x = start[0] + size if end[0] > start[0] else start[0] - size
|
||||||
|
end_y = start[1] + size if end[1] > start[1] else start[1] - size
|
||||||
|
return {"type": "rectangle", "start": start, "end": (end_x, end_y)}
|
||||||
|
elif self.tool_mode == "oval":
|
||||||
|
center = ((start[0] + end[0]) // 2, (start[1] + end[1]) // 2)
|
||||||
|
axes = (abs(end[0] - start[0]) // 2, abs(end[1] - start[1]) // 2)
|
||||||
|
return {"type": "ellipse", "center": center, "axes": axes}
|
||||||
|
elif self.tool_mode == "triangle":
|
||||||
|
return {"type": "polygon", "points": self._get_triangle_points(start, end)}
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _draw_shape_preview(self, start, end):
|
||||||
|
overlay = self.display_image.copy()
|
||||||
|
color = (0, 0, 255) # Red preview
|
||||||
|
|
||||||
|
params = self._get_shape_params(start, end)
|
||||||
|
if not params:
|
||||||
|
return
|
||||||
|
|
||||||
|
if params["type"] == "rectangle":
|
||||||
|
cv2.rectangle(overlay, params["start"], params["end"], color, -1)
|
||||||
|
elif params["type"] == "ellipse" and params["axes"][0] > 0 and params["axes"][1] > 0:
|
||||||
|
cv2.ellipse(overlay, params["center"], params["axes"], 0, 0, 360, color, -1)
|
||||||
|
elif params["type"] == "polygon":
|
||||||
|
cv2.fillPoly(overlay, [params["points"]], color)
|
||||||
|
|
||||||
|
cv2.addWeighted(overlay, 0.3, self.display_image, 0.7, 0, self.display_image)
|
||||||
|
|
||||||
|
def _draw_shape_on_mask(self, start, end):
|
||||||
|
params = self._get_shape_params(start, end)
|
||||||
|
if not params:
|
||||||
|
return
|
||||||
|
|
||||||
|
if params["type"] == "rectangle":
|
||||||
|
cv2.rectangle(self.mask, params["start"], params["end"], 255, -1)
|
||||||
|
elif params["type"] == "ellipse" and params["axes"][0] > 0 and params["axes"][1] > 0:
|
||||||
|
cv2.ellipse(self.mask, params["center"], params["axes"], 0, 0, 360, 255, -1)
|
||||||
|
elif params["type"] == "polygon":
|
||||||
|
cv2.fillPoly(self.mask, [params["points"]], 255)
|
||||||
|
|
||||||
|
def _get_triangle_points(self, start, end):
|
||||||
|
# Create isosceles triangle
|
||||||
|
base_center = ((start[0] + end[0]) // 2, end[1])
|
||||||
|
apex = (base_center[0], start[1])
|
||||||
|
left = (start[0], end[1])
|
||||||
|
right = (end[0], end[1])
|
||||||
|
return np.array([apex, left, right], np.int32)
|
||||||
|
|
||||||
def update_display(self):
|
def update_display(self):
|
||||||
# Create a copy of the original image
|
# Create a copy of the original image
|
||||||
self.display_image = self.original_image.copy()
|
self.display_image = self.original_image.copy()
|
||||||
@ -75,8 +183,20 @@ class MaskCreator:
|
|||||||
alpha = 0.5 # Transparency level
|
alpha = 0.5 # Transparency level
|
||||||
cv2.addWeighted(self.overlay, alpha, self.display_image, 1 - alpha, 0, self.display_image)
|
cv2.addWeighted(self.overlay, alpha, self.display_image, 1 - alpha, 0, self.display_image)
|
||||||
|
|
||||||
# Draw brush size indicator in the corner
|
# Draw tool info in the corner
|
||||||
text = f"Brush Size: {self.brush_size} (Hotkeys 1-9: change brush size) "
|
mode_display = {
|
||||||
|
"brush": "Brush",
|
||||||
|
"erase": "Erase",
|
||||||
|
"rectangle": "Rectangle",
|
||||||
|
"square": "Square",
|
||||||
|
"oval": "Oval",
|
||||||
|
"triangle": "Triangle",
|
||||||
|
}
|
||||||
|
text = f"Tool: {mode_display.get(self.tool_mode, self.tool_mode)} | "
|
||||||
|
if self.tool_mode in ["brush", "erase"]:
|
||||||
|
text += f"Size: {self.brush_size} (1-9) | "
|
||||||
|
text += f"History: [{len(self.mask_history)}]"
|
||||||
|
|
||||||
cv2.putText(
|
cv2.putText(
|
||||||
self.display_image,
|
self.display_image,
|
||||||
text,
|
text,
|
||||||
@ -96,10 +216,6 @@ class MaskCreator:
|
|||||||
cv2.imwrite(output_path, self.mask)
|
cv2.imwrite(output_path, self.mask)
|
||||||
print(f"Mask saved to {output_path}")
|
print(f"Mask saved to {output_path}")
|
||||||
|
|
||||||
def reset_mask(self):
|
|
||||||
self.mask = np.zeros(self.original_image.shape[:2], dtype=np.uint8)
|
|
||||||
self.update_display()
|
|
||||||
|
|
||||||
def set_brush_size(self, size_key):
|
def set_brush_size(self, size_key):
|
||||||
if size_key in self.BRUSH_SIZES:
|
if size_key in self.BRUSH_SIZES:
|
||||||
self.brush_size = self.BRUSH_SIZES[size_key]
|
self.brush_size = self.BRUSH_SIZES[size_key]
|
||||||
@ -114,21 +230,46 @@ class MaskCreator:
|
|||||||
key = cv2.waitKey(1) & 0xFF
|
key = cv2.waitKey(1) & 0xFF
|
||||||
key_char = chr(key) if key < 128 else ""
|
key_char = chr(key) if key < 128 else ""
|
||||||
|
|
||||||
# Check for brush size hotkeys (1-5)
|
# Check for brush size hotkeys (1-9)
|
||||||
if key_char in self.BRUSH_SIZES:
|
if key_char in self.BRUSH_SIZES and self.tool_mode in ["brush", "erase"]:
|
||||||
self.set_brush_size(key_char)
|
self.set_brush_size(key_char)
|
||||||
|
|
||||||
|
# Tool selection
|
||||||
|
elif key == ord("b"):
|
||||||
|
self.tool_mode = "brush"
|
||||||
|
print("Switched to brush tool")
|
||||||
|
self.update_display()
|
||||||
|
elif key == ord("e"):
|
||||||
|
self.tool_mode = "erase"
|
||||||
|
print("Switched to erase tool")
|
||||||
|
self.update_display()
|
||||||
|
elif key == ord("r"):
|
||||||
|
self.tool_mode = "rectangle"
|
||||||
|
print("Switched to rectangle tool")
|
||||||
|
self.update_display()
|
||||||
|
elif key == ord("q"):
|
||||||
|
self.tool_mode = "square"
|
||||||
|
print("Switched to square tool")
|
||||||
|
self.update_display()
|
||||||
|
elif key == ord("o"):
|
||||||
|
self.tool_mode = "oval"
|
||||||
|
print("Switched to oval tool")
|
||||||
|
self.update_display()
|
||||||
|
elif key == ord("t"):
|
||||||
|
self.tool_mode = "triangle"
|
||||||
|
print("Switched to triangle tool")
|
||||||
|
self.update_display()
|
||||||
|
|
||||||
|
# Undo (press 'u')
|
||||||
|
elif key == ord("u"):
|
||||||
|
self.undo()
|
||||||
|
|
||||||
# Save mask (press 's')
|
# Save mask (press 's')
|
||||||
elif key == ord("s"):
|
elif key == ord("s"):
|
||||||
self.save_mask(self.mask_output_path)
|
self.save_mask(self.mask_output_path)
|
||||||
|
|
||||||
# Reset mask (press 'r')
|
# Quit (press ESC)
|
||||||
elif key == ord("r"):
|
elif key == 27:
|
||||||
self.reset_mask()
|
|
||||||
print("Mask reset")
|
|
||||||
|
|
||||||
# Quit (press 'q' or ESC)
|
|
||||||
elif key == ord("q") or key == 27:
|
|
||||||
break
|
break
|
||||||
|
|
||||||
cv2.destroyAllWindows()
|
cv2.destroyAllWindows()
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user