Nisha56 commited on
Commit
28470ee
·
1 Parent(s): cfb9834

Week 3 complete: canvas, drag, resize, delete, export, history

Browse files
Files changed (6) hide show
  1. config.py +3 -2
  2. grounding.py +0 -2
  3. layer_extractor.py +102 -31
  4. main.py +37 -9
  5. pipeline.py +66 -61
  6. segment.py +40 -191
config.py CHANGED
@@ -4,8 +4,9 @@ from dotenv import load_dotenv
4
  load_dotenv()
5
 
6
  PROJECT_NAME = os.getenv("PROJECT_NAME", "poster-editor")
7
- UPLOAD_DIR = os.getenv("UPLOAD_DIR", "uploads")
8
- OUTPUT_DIR = os.getenv("OUTPUT_DIR", "outputs")
 
9
  MAX_FILE_SIZE_MB = int(os.getenv("MAX_FILE_SIZE_MB", 20))
10
 
11
  os.makedirs(UPLOAD_DIR, exist_ok=True)
 
4
  load_dotenv()
5
 
6
  PROJECT_NAME = os.getenv("PROJECT_NAME", "poster-editor")
7
+ BACKEND_DIR = os.path.dirname(os.path.abspath(__file__))
8
+ UPLOAD_DIR = os.path.join(BACKEND_DIR, os.getenv("UPLOAD_DIR", "uploads"))
9
+ OUTPUT_DIR = os.path.join(BACKEND_DIR, os.getenv("OUTPUT_DIR", "outputs"))
10
  MAX_FILE_SIZE_MB = int(os.getenv("MAX_FILE_SIZE_MB", 20))
11
 
12
  os.makedirs(UPLOAD_DIR, exist_ok=True)
grounding.py CHANGED
@@ -52,8 +52,6 @@ def detect_objects(image_path: str, model, text_prompt: str = None) -> list:
52
  # load_image returns (PIL image, transformed tensor) — both needed
53
  image_pil, image_tensor = load_image(image_path)
54
 
55
- print(type(image_pil))
56
- print(image_pil.shape)
57
 
58
  img_h, img_w = image_pil.shape[:2]
59
 
 
52
  # load_image returns (PIL image, transformed tensor) — both needed
53
  image_pil, image_tensor = load_image(image_path)
54
 
 
 
55
 
56
  img_h, img_w = image_pil.shape[:2]
57
 
layer_extractor.py CHANGED
@@ -1,49 +1,120 @@
1
  import os
 
2
  import json
 
 
 
 
 
 
 
3
  from grounding import load_grounding_model, detect_objects
4
  from ocr import extract_text
5
 
6
  BACKEND_DIR = os.path.dirname(os.path.abspath(__file__))
7
-
8
- # labels that GroundingDINO detects but we handle separately or don't need
9
  SKIP_LABELS = {"background", "text", "button"}
10
 
11
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
12
  def build_layers(image_path: str) -> list:
13
  """
14
- Run GroundingDINO + OCR on image, combine results into a flat layer list.
15
- Each layer has: id, type, x, y, w, h, confidence (+ text for text layers).
16
- This JSON is what the Fabric.js canvas loads to render editable elements.
 
 
 
 
17
  """
18
  layers = []
19
  layer_id = 1
20
 
21
- # --- object layers via GroundingDINO ---
22
- print("[LAYERS] Running object detection...")
 
 
 
 
 
23
  gdino_model = load_grounding_model()
24
  detections = detect_objects(image_path, gdino_model)
25
 
26
- for det in detections:
27
- if det["label"].lower() in SKIP_LABELS:
28
- continue
29
-
30
- layers.append({
31
- "id": layer_id,
32
- "type": "object",
33
- "label": det["label"],
34
- "x": det["x1"],
35
- "y": det["y1"],
36
- "w": det["w"],
37
- "h": det["h"],
38
- "confidence": det["confidence"],
39
- })
40
- layer_id += 1
41
-
42
- # --- text layers via PaddleOCR ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
43
  print("[LAYERS] Running OCR...")
44
  text_blocks = extract_text(image_path)
45
 
46
  for block in text_blocks:
 
 
 
 
 
47
  layers.append({
48
  "id": layer_id,
49
  "type": "text",
@@ -53,10 +124,12 @@ def build_layers(image_path: str) -> list:
53
  "w": block["w"],
54
  "h": block["h"],
55
  "confidence": block["confidence"],
 
 
56
  })
57
  layer_id += 1
58
 
59
- print(f"[LAYERS] Built {len(layers)} layers ({len(detections)} objects + {len(text_blocks)} text)")
60
  return layers
61
 
62
 
@@ -81,20 +154,18 @@ if __name__ == "__main__":
81
  exit(1)
82
 
83
  print("=" * 50)
84
- print("Building Editable Layers")
85
  print("=" * 50)
86
 
87
  layers = build_layers(TEST_IMAGE)
88
 
89
  print(f"\nGenerated {len(layers)} layers:")
90
  for layer in layers:
 
91
  if layer["type"] == "text":
92
- print(f" [{layer['id']}] TEXT '{layer['text']}' at ({layer['x']},{layer['y']})")
93
  else:
94
- print(f" [{layer['id']}] OBJECT '{layer['label']}' at ({layer['x']},{layer['y']})")
95
 
96
  save_layers(layers)
97
-
98
- print("=" * 50)
99
- print("Open outputs/layers.json")
100
  print("=" * 50)
 
1
  import os
2
+ import sys
3
  import json
4
+ import base64
5
+ import io
6
+ import cv2
7
+ import torch
8
+ import numpy as np
9
+ from PIL import Image
10
+
11
  from grounding import load_grounding_model, detect_objects
12
  from ocr import extract_text
13
 
14
  BACKEND_DIR = os.path.dirname(os.path.abspath(__file__))
 
 
15
  SKIP_LABELS = {"background", "text", "button"}
16
 
17
 
18
+ def pil_to_base64(img: Image.Image, fmt: str = "PNG") -> str:
19
+ """Convert PIL image to base64 string."""
20
+ buffer = io.BytesIO()
21
+ img.save(buffer, format=fmt)
22
+ return base64.b64encode(buffer.getvalue()).decode("utf-8")
23
+
24
+
25
+ def crop_text_layer(image_path: str, x: int, y: int, w: int, h: int) -> str:
26
+ """
27
+ Crop text region as transparent PNG.
28
+ Text layers use simple rectangular crop with white pixels made transparent.
29
+ """
30
+ img = Image.open(image_path).convert("RGBA")
31
+ crop = img.crop((x, y, x + w, y + h))
32
+ return pil_to_base64(crop, "PNG")
33
+
34
+
35
  def build_layers(image_path: str) -> list:
36
  """
37
+ Full layer extraction pipeline:
38
+ 1. GroundingDINO finds named objects with bounding boxes
39
+ 2. SAM2 refines each box into a precise pixel mask
40
+ 3. PaddleOCR finds all text blocks
41
+ 4. Each layer gets a transparent PNG crop
42
+
43
+ Returns flat list of layers ready for Fabric.js canvas.
44
  """
45
  layers = []
46
  layer_id = 1
47
 
48
+ # --- load image ---
49
+ image_bgr = cv2.imread(image_path)
50
+ image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)
51
+ img_h, img_w = image_rgb.shape[:2]
52
+
53
+ # --- step 1: GroundingDINO object detection ---
54
+ print("[LAYERS] Running GroundingDINO...")
55
  gdino_model = load_grounding_model()
56
  detections = detect_objects(image_path, gdino_model)
57
 
58
+ # filter background/text labels
59
+ obj_detections = [d for d in detections if d["label"].lower() not in SKIP_LABELS]
60
+ print(f"[LAYERS] {len(obj_detections)} objects to segment with SAM2")
61
+
62
+ # --- step 2: SAM2 precise masks for each object ---
63
+ if obj_detections:
64
+ from segment import load_sam2_model, get_mask_for_box, mask_to_transparent_png
65
+ predictor = load_sam2_model()
66
+
67
+ for det in obj_detections:
68
+ box = [det["x1"], det["y1"], det["x2"], det["y2"]]
69
+
70
+ try:
71
+ mask = get_mask_for_box(predictor, image_rgb, box)
72
+
73
+ # get tight bounding box from mask
74
+ rows = np.where(mask.any(axis=1))[0]
75
+ cols = np.where(mask.any(axis=0))[0]
76
+ if len(rows) == 0: continue
77
+
78
+ y1, y2 = int(rows.min()), int(rows.max())
79
+ x1, x2 = int(cols.min()), int(cols.max())
80
+
81
+ # create transparent PNG
82
+ png_img = mask_to_transparent_png(image_rgb, mask)
83
+ b64 = pil_to_base64(png_img, "PNG")
84
+
85
+ layers.append({
86
+ "id": layer_id,
87
+ "type": "object",
88
+ "label": det["label"],
89
+ "x": x1,
90
+ "y": y1,
91
+ "w": x2 - x1,
92
+ "h": y2 - y1,
93
+ "confidence": det["confidence"],
94
+ "base64": b64,
95
+ "format": "png",
96
+ })
97
+ layer_id += 1
98
+
99
+ except Exception as e:
100
+ print(f"[LAYERS] SAM2 failed for {det['label']}: {e}")
101
+ continue
102
+
103
+ # free SAM2 VRAM before OCR
104
+ del predictor
105
+ torch.cuda.empty_cache() if __import__('torch').cuda.is_available() else None
106
+ print("[LAYERS] SAM2 VRAM freed")
107
+
108
+ # --- step 3: OCR text layers ---
109
  print("[LAYERS] Running OCR...")
110
  text_blocks = extract_text(image_path)
111
 
112
  for block in text_blocks:
113
+ b64 = crop_text_layer(
114
+ image_path,
115
+ block["x"], block["y"],
116
+ block["w"], block["h"]
117
+ )
118
  layers.append({
119
  "id": layer_id,
120
  "type": "text",
 
124
  "w": block["w"],
125
  "h": block["h"],
126
  "confidence": block["confidence"],
127
+ "base64": b64,
128
+ "format": "png",
129
  })
130
  layer_id += 1
131
 
132
+ print(f"[LAYERS] Built {len(layers)} layers ({len(obj_detections)} objects + {len(text_blocks)} text)")
133
  return layers
134
 
135
 
 
154
  exit(1)
155
 
156
  print("=" * 50)
157
+ print("Building Transparent Layers")
158
  print("=" * 50)
159
 
160
  layers = build_layers(TEST_IMAGE)
161
 
162
  print(f"\nGenerated {len(layers)} layers:")
163
  for layer in layers:
164
+ fmt = layer.get("format", "jpg")
165
  if layer["type"] == "text":
166
+ print(f" [{layer['id']}] TEXT '{layer['text']}' at ({layer['x']},{layer['y']}) [{fmt}]")
167
  else:
168
+ print(f" [{layer['id']}] OBJECT '{layer['label']}' at ({layer['x']},{layer['y']}) [{fmt}]")
169
 
170
  save_layers(layers)
 
 
 
171
  print("=" * 50)
main.py CHANGED
@@ -76,14 +76,15 @@ async def upload_image(
76
  raise HTTPException(500, f"Pipeline failed: {str(e)}")
77
 
78
  return JSONResponse({
79
- "project_id": project.id,
80
- "file_id": file_id,
81
- "filename": file.filename,
82
- "image_w": result["image_w"],
83
- "image_h": result["image_h"],
84
- "image_base64": result["image_base64"],
85
- "layers": result["layers"],
86
- "processing_time_s": result["processing_time_s"],
 
87
  })
88
 
89
 
@@ -166,4 +167,31 @@ async def inpaint_element(request: dict):
166
 
167
  except Exception as e:
168
  unload_inpaint_model()
169
- raise HTTPException(500, f"Inpainting failed: {str(e)}")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
76
  raise HTTPException(500, f"Pipeline failed: {str(e)}")
77
 
78
  return JSONResponse({
79
+ "project_id": project.id,
80
+ "file_id": file_id,
81
+ "filename": file.filename,
82
+ "image_w": result["image_w"],
83
+ "image_h": result["image_h"],
84
+ "background_base64": result["background_base64"],
85
+ "original_base64": result["original_base64"],
86
+ "layers": result["layers"],
87
+ "processing_time_s": result["processing_time_s"],
88
  })
89
 
90
 
 
167
 
168
  except Exception as e:
169
  unload_inpaint_model()
170
+ raise HTTPException(500, f"Inpainting failed: {str(e)}")
171
+
172
+
173
+ @app.get("/projects/{file_id}/layers")
174
+ async def get_project_layers(file_id: str, db: Session = Depends(get_db)):
175
+ """Re-run pipeline on existing upload to get full layer data."""
176
+ project = get_project_by_file_id(db, file_id)
177
+ if not project:
178
+ raise HTTPException(404, "Project not found")
179
+
180
+ if not os.path.exists(project.upload_path):
181
+ raise HTTPException(404, "Original file no longer exists")
182
+
183
+ try:
184
+ result = run_pipeline(project.upload_path)
185
+ return JSONResponse({
186
+ "project_id": project.id,
187
+ "file_id": file_id,
188
+ "filename": project.filename,
189
+ "image_w": result["image_w"],
190
+ "image_h": result["image_h"],
191
+ "background_base64": result["background_base64"],
192
+ "original_base64": result["original_base64"],
193
+ "layers": result["layers"],
194
+ "processing_time_s": result["processing_time_s"],
195
+ })
196
+ except Exception as e:
197
+ raise HTTPException(500, f"Pipeline failed: {str(e)}")
pipeline.py CHANGED
@@ -2,98 +2,105 @@ import os
2
  import time
3
  import base64
4
  import json
 
 
 
5
  from PIL import Image
6
 
7
- from layer_extractor import build_layers, save_layers
8
 
9
  BACKEND_DIR = os.path.dirname(os.path.abspath(__file__))
10
 
11
 
12
  def image_to_base64(image_path: str) -> str:
13
- """Read image file and return base64 string for sending to frontend."""
14
  with open(image_path, "rb") as f:
15
  return base64.b64encode(f.read()).decode("utf-8")
16
 
17
 
18
- def crop_layer_image(image_path: str, x: int, y: int, w: int, h: int) -> str:
 
 
 
 
 
 
19
  """
20
- Crop a region from the original image and return as base64.
21
- This gives the frontend the actual pixel content of each layer
22
- so Fabric.js can render it on the canvas.
23
  """
24
- img = Image.open(image_path).convert("RGB")
25
- region = img.crop((x, y, x + w, y + h))
 
 
 
26
 
27
- # save cropped region to temp file then encode
28
- temp_path = os.path.join(BACKEND_DIR, "outputs", f"crop_temp.jpg")
29
- region.save(temp_path, quality=95)
 
 
 
 
 
 
 
 
 
 
 
30
 
31
- return image_to_base64(temp_path)
 
 
32
 
33
 
34
  def run_pipeline(image_path: str) -> dict:
35
  """
36
- Full pipeline: image → layer JSON ready for Fabric.js canvas.
37
-
38
- Steps:
39
- 1. GroundingDINO detects named objects
40
- 2. PaddleOCR reads all text
41
- 3. Each layer gets its cropped image as base64
42
- 4. Returns complete response the frontend expects
43
-
44
- Returns:
45
- {
46
- image_w, image_h,
47
- image_base64, ← full poster for canvas background
48
- layers: [
49
- {id, type, label/text, x, y, w, h, confidence, base64}
50
- ]
51
- }
52
  """
53
  start = time.time()
54
- print(f"\n[PIPELINE] Starting pipeline for: {os.path.basename(image_path)}")
55
 
56
- # get image dimensions
57
- img = Image.open(image_path)
58
  img_w, img_h = img.size
59
- print(f"[PIPELINE] Image size: {img_w}x{img_h}")
60
 
61
- # run detection + OCR
62
  layers = build_layers(image_path)
 
63
 
64
- # attach cropped base64 image to each layer
65
- # frontend needs this to render each element on canvas
66
- print("[PIPELINE] Cropping layer images...")
67
- for layer in layers:
68
- try:
69
- layer["base64"] = crop_layer_image(
70
- image_path,
71
- layer["x"], layer["y"],
72
- layer["w"], layer["h"]
73
- )
74
- except Exception as e:
75
- print(f"[PIPELINE] Crop failed for layer {layer['id']}: {e}")
76
- layer["base64"] = ""
77
 
78
  elapsed = round(time.time() - start, 2)
79
- print(f"[PIPELINE] Done in {elapsed}s — {len(layers)} layers extracted")
80
 
81
  return {
82
- "image_w": img_w,
83
- "image_h": img_h,
84
- "image_base64": image_to_base64(image_path),
85
- "layers": layers,
86
- "processing_time_s": elapsed,
 
87
  }
88
 
89
 
90
  def save_pipeline_output(result: dict, output_dir: str = None) -> str:
91
- """Save pipeline result to JSON — useful for debugging and development."""
92
  if output_dir is None:
93
  output_dir = os.path.join(BACKEND_DIR, "outputs")
94
  os.makedirs(output_dir, exist_ok=True)
95
 
96
- # save without base64 blobs — they're huge and unreadable in JSON viewer
97
  result_slim = {
98
  "image_w": result["image_w"],
99
  "image_h": result["image_h"],
@@ -109,7 +116,7 @@ def save_pipeline_output(result: dict, output_dir: str = None) -> str:
109
  with open(out_path, "w", encoding="utf-8") as f:
110
  json.dump(result_slim, f, indent=2, ensure_ascii=False)
111
 
112
- print(f"[PIPELINE] Saved slim output → {out_path}")
113
  return out_path
114
 
115
 
@@ -121,16 +128,14 @@ if __name__ == "__main__":
121
  exit(1)
122
 
123
  print("=" * 50)
124
- print("Running Full Editify Pipeline")
125
  print("=" * 50)
126
 
127
  result = run_pipeline(TEST_IMAGE)
128
  save_pipeline_output(result)
129
 
130
  print("=" * 50)
131
- print(f"Pipeline complete.")
132
- print(f" Layers: {len(result['layers'])}")
133
- print(f" Image size: {result['image_w']}x{result['image_h']}")
134
- print(f" Processing time: {result['processing_time_s']}s")
135
- print(f" Image base64: {len(result['image_base64'])} chars")
136
  print("=" * 50)
 
2
  import time
3
  import base64
4
  import json
5
+ import io
6
+ import cv2
7
+ import numpy as np
8
  from PIL import Image
9
 
10
+ from layer_extractor import build_layers
11
 
12
  BACKEND_DIR = os.path.dirname(os.path.abspath(__file__))
13
 
14
 
15
  def image_to_base64(image_path: str) -> str:
 
16
  with open(image_path, "rb") as f:
17
  return base64.b64encode(f.read()).decode("utf-8")
18
 
19
 
20
+ def pil_to_base64(img: Image.Image, fmt: str = "PNG") -> str:
21
+ buffer = io.BytesIO()
22
+ img.save(buffer, format=fmt)
23
+ return base64.b64encode(buffer.getvalue()).decode("utf-8")
24
+
25
+
26
+ def reconstruct_background(image_path: str, layers: list) -> str:
27
  """
28
+ Build ONE combined mask covering all detected layers.
29
+ Run a single OpenCV TELEA inpaint to reconstruct the clean background.
30
+ Returns base64 JPEG of the clean background.
31
  """
32
+ image_bgr = cv2.imread(image_path)
33
+ h, w = image_bgr.shape[:2]
34
+
35
+ # combined mask — white where any layer exists
36
+ combined_mask = np.zeros((h, w), dtype=np.uint8)
37
 
38
+ for layer in layers:
39
+ x = max(0, layer["x"])
40
+ y = max(0, layer["y"])
41
+ x2 = min(w, layer["x"] + layer["w"])
42
+ y2 = min(h, layer["y"] + layer["h"])
43
+ combined_mask[y:y2, x:x2] = 255
44
+
45
+ # single TELEA inpaint — fast, works for MVP, replaceable later
46
+ clean_bg = cv2.inpaint(
47
+ image_bgr,
48
+ combined_mask,
49
+ inpaintRadius=25,
50
+ flags=cv2.INPAINT_TELEA
51
+ )
52
 
53
+ # encode to base64
54
+ _, buffer = cv2.imencode('.jpg', clean_bg, [cv2.IMWRITE_JPEG_QUALITY, 95])
55
+ return base64.b64encode(buffer).decode("utf-8")
56
 
57
 
58
  def run_pipeline(image_path: str) -> dict:
59
  """
60
+ Full pipeline:
61
+ 1. Extract all layers as transparent PNGs
62
+ 2. Build combined mask of all layer regions
63
+ 3. Reconstruct clean background in ONE inpaint pass
64
+ 4. Return clean background + independent layers
65
+
66
+ Editor renders: clean background + layer PNGs
67
+ Every element exists exactly once — no duplicates possible.
68
+ Editing is instant — no AI calls needed during editing.
 
 
 
 
 
 
 
69
  """
70
  start = time.time()
71
+ print(f"\n[PIPELINE] Starting: {os.path.basename(image_path)}")
72
 
73
+ img = Image.open(image_path)
 
74
  img_w, img_h = img.size
75
+ print(f"[PIPELINE] Image: {img_w}x{img_h}")
76
 
77
+ # step 1 — extract all layers
78
  layers = build_layers(image_path)
79
+ print(f"[PIPELINE] Extracted {len(layers)} layers")
80
 
81
+ # step 2 — reconstruct background once using combined mask
82
+ print("[PIPELINE] Reconstructing clean background...")
83
+ bg_base64 = reconstruct_background(image_path, layers)
84
+ print("[PIPELINE] Background reconstructed")
 
 
 
 
 
 
 
 
 
85
 
86
  elapsed = round(time.time() - start, 2)
87
+ print(f"[PIPELINE] Done in {elapsed}s")
88
 
89
  return {
90
+ "image_w": img_w,
91
+ "image_h": img_h,
92
+ "background_base64": bg_base64, # clean background, no elements
93
+ "original_base64": image_to_base64(image_path), # for inpainting reference
94
+ "layers": layers,
95
+ "processing_time_s": elapsed,
96
  }
97
 
98
 
99
  def save_pipeline_output(result: dict, output_dir: str = None) -> str:
 
100
  if output_dir is None:
101
  output_dir = os.path.join(BACKEND_DIR, "outputs")
102
  os.makedirs(output_dir, exist_ok=True)
103
 
 
104
  result_slim = {
105
  "image_w": result["image_w"],
106
  "image_h": result["image_h"],
 
116
  with open(out_path, "w", encoding="utf-8") as f:
117
  json.dump(result_slim, f, indent=2, ensure_ascii=False)
118
 
119
+ print(f"[PIPELINE] Saved → {out_path}")
120
  return out_path
121
 
122
 
 
128
  exit(1)
129
 
130
  print("=" * 50)
131
+ print("Running Editify Pipeline")
132
  print("=" * 50)
133
 
134
  result = run_pipeline(TEST_IMAGE)
135
  save_pipeline_output(result)
136
 
137
  print("=" * 50)
138
+ print(f"Layers: {len(result['layers'])}")
139
+ print(f"Image: {result['image_w']}x{result['image_h']}")
140
+ print(f"Time: {result['processing_time_s']}s")
 
 
141
  print("=" * 50)
segment.py CHANGED
@@ -1,27 +1,23 @@
1
- import torch
2
- import numpy as np
3
- import cv2
4
  import os
5
- import json
6
  import sys
 
 
 
7
 
8
- # tell Python where SAM2 code lives
9
  sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'sam2'))
10
 
11
  from sam2.build_sam import build_sam2
12
  from sam2.sam2_image_predictor import SAM2ImagePredictor
13
 
14
- # config
15
  MODEL_CFG = "configs/sam2.1/sam2.1_hiera_s.yaml"
16
  CHECKPOINT = os.path.join(os.path.dirname(__file__), "models", "sam2.1_hiera_small.pt")
17
  DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
18
- SAM2_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'sam2')
19
 
20
  print(f"[SAM2] Using device: {DEVICE}")
21
 
22
 
23
- def load_model():
24
- """Load SAM2 into GPU memory. Call once at server startup."""
25
  os.chdir(SAM2_DIR)
26
  model = build_sam2(MODEL_CFG, CHECKPOINT, device=DEVICE)
27
  predictor = SAM2ImagePredictor(model)
@@ -29,197 +25,50 @@ def load_model():
29
  return predictor
30
 
31
 
32
- def generate_grid_points(image_h, image_w, rows=8, cols=8):
33
  """
34
- Create a grid of points spread across the image.
35
- SAM2 needs input points as prompts to know where to look.
36
- A 5x5 grid = 25 points covering the entire poster evenly.
37
  """
38
- points = []
39
- for r in range(1, rows + 1):
40
- for c in range(1, cols + 1):
41
- x = int(image_w * c / (cols + 1))
42
- y = int(image_h * r / (rows + 1))
43
- points.append([x, y])
44
- return np.array(points, dtype=np.float32)
45
-
46
-
47
- def segment_image(image_path: str, predictor) -> dict:
48
- """
49
- Main segmentation function.
50
- Input : path to any poster image
51
- Output: dict with list of layers [{id, score, x, y, w, h, area, mask}]
52
- """
53
- # load image
54
- image_bgr = cv2.imread(image_path)
55
- if image_bgr is None:
56
- raise ValueError(f"Could not read image at {image_path}")
57
-
58
- image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)
59
- h, w = image_rgb.shape[:2]
60
- print(f"[SAM2] Image loaded: {w}x{h}px")
61
-
62
- # encode image into SAM2 — this is the heavy step (~3-5 seconds)
63
  predictor.set_image(image_rgb)
64
- print("[SAM2] Image encoded, running segmentation...")
65
 
66
- # generate grid prompts
67
- grid_points = generate_grid_points(h, w, rows=5, cols=5)
68
- labels = np.ones(len(grid_points), dtype=np.int32)
69
 
70
- # run SAM2
71
  with torch.inference_mode():
72
  with torch.autocast(device_type=DEVICE, dtype=torch.float16):
73
  masks, scores, _ = predictor.predict(
74
- point_coords=grid_points,
75
- point_labels=labels,
76
- multimask_output=True,
 
77
  )
78
 
79
- print(f"[SAM2] Generated {len(masks)} raw masks")
80
-
81
- # flatten masks and scores if shape is (N, H, W)
82
- if masks.ndim == 3:
83
- masks_list = [masks[i] for i in range(masks.shape[0])]
84
- scores_list = [scores[i] for i in range(scores.shape[0])]
85
- else:
86
- masks_list = list(masks)
87
- scores_list = list(scores)
88
-
89
- # filter and clean masks
90
- layers = []
91
- seen_areas = set()
92
-
93
- for i, (mask, score) in enumerate(zip(masks_list, scores_list)):
94
- score = float(score)
95
- if score < 0.6:
96
- continue
97
-
98
- # get bounding box
99
- rows_idx = np.where(mask.any(axis=1))[0]
100
- cols_idx = np.where(mask.any(axis=0))[0]
101
- if len(rows_idx) == 0 or len(cols_idx) == 0:
102
- continue
103
-
104
- y1, y2 = int(rows_idx.min()), int(rows_idx.max())
105
- x1, x2 = int(cols_idx.min()), int(cols_idx.max())
106
- area = int((x2 - x1) * (y2 - y1))
107
-
108
- # skip noise and full-image background
109
- min_area = int(w * h * 0.005)
110
- max_area = int(w * h * 0.95)
111
- if area < min_area or area > max_area:
112
- continue
113
-
114
- # deduplicate similar masks
115
- area_key = area // 1000
116
- if area_key in seen_areas:
117
- continue
118
- seen_areas.add(area_key)
119
-
120
- layers.append({
121
- "id": i,
122
- "score": round(score, 3),
123
- "x": x1,
124
- "y": y1,
125
- "w": x2 - x1,
126
- "h": y2 - y1,
127
- "area": area,
128
- "mask": mask.astype(bool),
129
- })
130
-
131
- print(f"[SAM2] {len(layers)} clean layers after filtering")
132
- return {"layers": layers, "image_size": {"w": w, "h": h}}
133
-
134
-
135
- def save_debug_output(image_path: str, result: dict, output_dir: str = None):
136
- if output_dir is None:
137
- output_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "outputs")
138
  """
139
- Save visual output so you can SEE what SAM2 detected.
140
- - masked_overlay.jpg : poster with coloured regions over each object
141
- - layers_data.json : layer positions and scores
142
  """
143
- os.makedirs(output_dir, exist_ok=True)
144
- image_bgr = cv2.imread(image_path)
145
- overlay = image_bgr.copy()
146
-
147
- colours = [
148
- (255, 80, 80), (80, 255, 80), (80, 80, 255),
149
- (255, 255, 80), (255, 80, 255), (80, 255, 255),
150
- (200, 140, 80), (140, 200, 80), (80, 140, 200),
151
- ]
152
-
153
- json_layers = []
154
-
155
- for idx, layer in enumerate(result["layers"]):
156
- colour = colours[idx % len(colours)]
157
- mask = layer["mask"]
158
-
159
- # semi-transparent fill
160
- overlay[mask] = (
161
- overlay[mask] * 0.45 + np.array(colour) * 0.55
162
- ).astype(np.uint8)
163
-
164
- # bounding box
165
- cv2.rectangle(
166
- overlay,
167
- (layer["x"], layer["y"]),
168
- (layer["x"] + layer["w"], layer["y"] + layer["h"]),
169
- colour, 2
170
- )
171
-
172
- # label
173
- cv2.putText(
174
- overlay,
175
- f"L{idx} {layer['score']}",
176
- (layer["x"] + 4, layer["y"] + 18),
177
- cv2.FONT_HERSHEY_SIMPLEX, 0.5, colour, 1
178
- )
179
-
180
- json_layers.append({
181
- "id": layer["id"],
182
- "score": layer["score"],
183
- "x": layer["x"],
184
- "y": layer["y"],
185
- "w": layer["w"],
186
- "h": layer["h"],
187
- "area": layer["area"],
188
- })
189
-
190
- out_img = os.path.join(output_dir, "masked_overlay.jpg")
191
- out_json = os.path.join(output_dir, "layers_data.json")
192
-
193
- cv2.imwrite(out_img, overlay)
194
- with open(out_json, "w") as f:
195
- json.dump({
196
- "layers": json_layers,
197
- "image_size": result["image_size"]
198
- }, f, indent=2)
199
-
200
- print(f"[SAM2] Saved overlay → {out_img}")
201
- print(f"[SAM2] Saved JSON → {out_json}")
202
- return out_img, out_json
203
-
204
-
205
- # run directly to test
206
- if __name__ == "__main__":
207
- BACKEND_DIR = os.path.dirname(os.path.abspath(__file__))
208
- TEST_IMAGE = os.path.join(BACKEND_DIR, "test_images", "sale_img.jpg")
209
-
210
- if not os.path.exists(TEST_IMAGE):
211
- print(f"ERROR: No image at {TEST_IMAGE}")
212
- exit(1)
213
-
214
- print("=" * 50)
215
- print("Testing SAM2 segmentation")
216
- print("=" * 50)
217
-
218
- predictor = load_model()
219
- result = segment_image(TEST_IMAGE, predictor)
220
- save_debug_output(TEST_IMAGE, result)
221
-
222
- print("=" * 50)
223
- print(f"DONE. Found {len(result['layers'])} layers.")
224
- print("Open outputs/masked_overlay.jpg to see results.")
225
- print("=" * 50)
 
 
 
 
1
  import os
 
2
  import sys
3
+ import torch
4
+ import numpy as np
5
+ from PIL import Image
6
 
 
7
  sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'sam2'))
8
 
9
  from sam2.build_sam import build_sam2
10
  from sam2.sam2_image_predictor import SAM2ImagePredictor
11
 
 
12
  MODEL_CFG = "configs/sam2.1/sam2.1_hiera_s.yaml"
13
  CHECKPOINT = os.path.join(os.path.dirname(__file__), "models", "sam2.1_hiera_small.pt")
14
  DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
15
+ SAM2_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'sam2')
16
 
17
  print(f"[SAM2] Using device: {DEVICE}")
18
 
19
 
20
+ def load_sam2_model():
 
21
  os.chdir(SAM2_DIR)
22
  model = build_sam2(MODEL_CFG, CHECKPOINT, device=DEVICE)
23
  predictor = SAM2ImagePredictor(model)
 
25
  return predictor
26
 
27
 
28
+ def get_mask_for_box(predictor, image_rgb: np.ndarray, box: list) -> np.ndarray:
29
  """
30
+ Get precise pixel mask for a single bounding box using SAM2.
31
+ box = [x1, y1, x2, y2] in absolute pixels.
32
+ Returns boolean mask same size as image.
33
  """
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
34
  predictor.set_image(image_rgb)
 
35
 
36
+ box_array = np.array(box, dtype=np.float32)
 
 
37
 
 
38
  with torch.inference_mode():
39
  with torch.autocast(device_type=DEVICE, dtype=torch.float16):
40
  masks, scores, _ = predictor.predict(
41
+ point_coords = None,
42
+ point_labels = None,
43
+ box = box_array[None, :], # SAM2 expects (1, 4)
44
+ multimask_output = True,
45
  )
46
 
47
+ # pick highest confidence mask
48
+ best_idx = scores.argmax()
49
+ best_mask = masks[best_idx].astype(bool)
50
+ return best_mask
51
+
52
+
53
+ def mask_to_transparent_png(image_rgb: np.ndarray, mask: np.ndarray) -> Image.Image:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
54
  """
55
+ Apply mask to image — keep masked pixels, make everything else transparent.
56
+ Returns RGBA PIL image.
 
57
  """
58
+ rgba = np.zeros((*image_rgb.shape[:2], 4), dtype=np.uint8)
59
+ rgba[..., :3] = image_rgb
60
+ rgba[..., 3] = (mask * 255).astype(np.uint8) # alpha = 255 where object is
61
+
62
+ # crop to bounding box of mask to minimise image size
63
+ rows = np.where(mask.any(axis=1))[0]
64
+ cols = np.where(mask.any(axis=0))[0]
65
+
66
+ if len(rows) == 0 or len(cols) == 0:
67
+ return Image.fromarray(rgba, 'RGBA')
68
+
69
+ y1, y2 = int(rows.min()), int(rows.max())
70
+ x1, x2 = int(cols.min()), int(cols.max())
71
+
72
+ cropped = rgba[y1:y2+1, x1:x2+1]
73
+ return Image.fromarray(cropped, 'RGBA')
74
+