multimodalart HF Staff commited on
Commit
121d73b
·
verified ·
1 Parent(s): 0b107ad

Upload folder using huggingface_hub

Browse files
Files changed (8) hide show
  1. .gitattributes +1 -0
  2. README.md +25 -7
  3. app.py +650 -0
  4. demo.jpg +3 -0
  5. demo_mask_0.jpg +0 -0
  6. demo_mask_1.jpg +0 -0
  7. demo_mask_2.jpg +0 -0
  8. requirements.txt +8 -0
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ demo.jpg filter=lfs diff=lfs merge=lfs -text
README.md CHANGED
@@ -1,13 +1,31 @@
1
  ---
2
- title: Perception Dlm
3
- emoji: 🌍
4
- colorFrom: gray
5
- colorTo: green
6
  sdk: gradio
7
  sdk_version: 6.19.0
8
- python_version: '3.12'
9
  app_file: app.py
10
- pinned: false
 
 
11
  ---
12
 
13
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: PerceptionDLM Region Captioning
3
+ emoji: 🎯
4
+ colorFrom: indigo
5
+ colorTo: yellow
6
  sdk: gradio
7
  sdk_version: 6.19.0
 
8
  app_file: app.py
9
+ short_description: Parallel region captioning with multimodal diffusion LLM
10
+ python_version: "3.12"
11
+ startup_duration_timeout: 1h
12
  ---
13
 
14
+ # PerceptionDLM Region Captioning
15
+
16
+ A Gradio demo for [MSALab/PerceptionDLM](https://huggingface.co/MSALab/PerceptionDLM), a 9.2B parameter multimodal diffusion language model for parallel region captioning.
17
+
18
+ ## How it works
19
+
20
+ Upload an image and one or more binary mask images. The model generates descriptions for all masked regions **simultaneously** in a single denoising process — avoiding the linear latency growth of autoregressive region captioners.
21
+
22
+ The decoding animation replays each diffusion step so you can watch captions emerge token by token.
23
+
24
+ ## Model details
25
+
26
+ - **Base:** LLaDA-8B (diffusion language model) + SigLIP2 vision encoder
27
+ - **Precision:** bfloat16
28
+ - **Region prompts:** up to 6 per image
29
+ - **Default inference:** 32 diffusion steps, generation length 32 per mask
30
+ - **Paper:** [arXiv:2606.19534](https://arxiv.org/abs/2606.19534)
31
+ - **Code:** [GitHub](https://github.com/MSALab-PKU/PerceptionDLM)
app.py ADDED
@@ -0,0 +1,650 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Gradio demo for PerceptionDLM parallel region captioning.
2
+
3
+ This app runs the PerceptionDLM model on ZeroGPU. Users upload an image and
4
+ one or more binary masks, and the model generates captions for all masked
5
+ regions in parallel via a single denoising process. The decoding animation
6
+ replays each diffusion step so you can watch captions emerge token by token.
7
+ """
8
+
9
+ import os
10
+
11
+ os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
12
+
13
+ import html as html_lib
14
+ import time
15
+ import random
16
+ from typing import Dict, List, Tuple
17
+
18
+ import spaces # MUST come before torch / any CUDA-touching import
19
+ import torch
20
+ import numpy as np
21
+ from PIL import Image
22
+
23
+ import gradio as gr
24
+ from transformers import AutoModel, AutoProcessor
25
+
26
+ # ---------------------------------------------------------------------------
27
+ # Model loading at module scope (ZeroGPU intercepts .to("cuda"))
28
+ # ---------------------------------------------------------------------------
29
+ MODEL_ID = "MSALab/PerceptionDLM"
30
+ DTYPE = torch.bfloat16
31
+
32
+ print(f"Loading processor from {MODEL_ID} ...")
33
+ PROCESSOR = AutoProcessor.from_pretrained(MODEL_ID, trust_remote_code=True)
34
+ TOKENIZER = PROCESSOR.tokenizer
35
+
36
+ print(f"Loading model from {MODEL_ID} ...")
37
+ MODEL = AutoModel.from_pretrained(
38
+ MODEL_ID,
39
+ torch_dtype=DTYPE,
40
+ trust_remote_code=True,
41
+ attn_implementation="sdpa",
42
+ )
43
+ MODEL.processor = PROCESSOR
44
+ MODEL.to("cuda")
45
+ MODEL.eval()
46
+ print("Model loaded.")
47
+
48
+ # ---------------------------------------------------------------------------
49
+ # Constants
50
+ # ---------------------------------------------------------------------------
51
+ MASK_ID = 126336 # token id for the LLaDA diffusion backbone
52
+ MASK_PLACEHOLDER = "\ue000" # sentinel for not-yet-revealed tokens
53
+ DEFAULT_PROMPT = "Describe each masked region in detail."
54
+
55
+ OVERLAY_COLORS = [
56
+ (239, 68, 68), # red
57
+ (16, 185, 129), # green
58
+ (59, 130, 246), # blue
59
+ (245, 158, 11), # amber
60
+ (236, 72, 153), # pink
61
+ (139, 92, 246), # violet
62
+ (6, 182, 212), # cyan
63
+ (132, 204, 22), # lime
64
+ ]
65
+
66
+ # ---------------------------------------------------------------------------
67
+ # Preprocessing helpers (adapted from demo/infer_pdmllm.py)
68
+ # ---------------------------------------------------------------------------
69
+
70
+ def find_closest_aspect_ratio(aspect_ratio, target_ratios, width, height, image_size):
71
+ best_ratio_diff = float('inf')
72
+ best_ratio = (1, 1)
73
+ area = width * height
74
+ for ratio in target_ratios:
75
+ target_aspect_ratio = ratio[0] / ratio[1]
76
+ ratio_diff = abs(aspect_ratio - target_aspect_ratio)
77
+ if ratio_diff < best_ratio_diff:
78
+ best_ratio_diff = ratio_diff
79
+ best_ratio = ratio
80
+ elif ratio_diff == best_ratio_diff:
81
+ if area > 0.5 * image_size * image_size * ratio[0] * ratio[1]:
82
+ best_ratio = ratio
83
+ return best_ratio
84
+
85
+
86
+ def dynamic_preprocess(image, min_num=1, max_num=6, image_size=512, use_thumbnail=True):
87
+ orig_width, orig_height = image.size
88
+ aspect_ratio = orig_width / orig_height
89
+ target_ratios = set(
90
+ (i, j) for n in range(min_num, max_num + 1)
91
+ for i in range(1, n + 1) for j in range(1, n + 1)
92
+ if i * j <= max_num and i * j >= min_num
93
+ )
94
+ target_ratios = sorted(target_ratios, key=lambda x: x[0] * x[1])
95
+ target_aspect_ratio = find_closest_aspect_ratio(
96
+ aspect_ratio, target_ratios, orig_width, orig_height, image_size
97
+ )
98
+ target_width = image_size * target_aspect_ratio[0]
99
+ target_height = image_size * target_aspect_ratio[1]
100
+ blocks = target_aspect_ratio[0] * target_aspect_ratio[1]
101
+ resized_img = image.resize((target_width, target_height))
102
+ processed_images = []
103
+ for i in range(blocks):
104
+ box = (
105
+ (i % (target_width // image_size)) * image_size,
106
+ (i // (target_width // image_size)) * image_size,
107
+ ((i % (target_width // image_size)) + 1) * image_size,
108
+ ((i // (target_width // image_size)) + 1) * image_size,
109
+ )
110
+ split_img = resized_img.crop(box)
111
+ processed_images.append(split_img)
112
+ assert len(processed_images) == blocks
113
+ if use_thumbnail and len(processed_images) != 1:
114
+ thumbnail_img = image.resize((image_size, image_size))
115
+ processed_images.append(thumbnail_img)
116
+ return processed_images
117
+
118
+
119
+ def sort_masks_by_area(masks: List[np.ndarray]):
120
+ areas = [np.sum(m) for m in masks]
121
+ return np.argsort(np.array(areas))[::-1]
122
+
123
+
124
+ def build_visual_prompt_matrices(
125
+ masks: List[np.ndarray],
126
+ prompt_numbers: int,
127
+ ) -> tuple:
128
+ if len(masks) > prompt_numbers:
129
+ raise ValueError(
130
+ f"Number of masks ({len(masks)}) exceeds prompt_numbers ({prompt_numbers})."
131
+ )
132
+ height, width = masks[0].shape
133
+ prompt_indexes = list(range(prompt_numbers))
134
+ selected_prompt_indexes = prompt_indexes[:len(masks)]
135
+ selected_prompt_tokens = [f"<Prompt{i}>" for i in selected_prompt_indexes]
136
+
137
+ filled_matrices = []
138
+ for prompt_id, mask in zip(selected_prompt_indexes, masks):
139
+ filled_matrix = np.full((height, width), 255, dtype=np.uint8)
140
+ fill_area = (filled_matrix == 255) & mask.astype(bool)
141
+ filled_matrix[fill_area] = prompt_id
142
+ filled_matrices.append(filled_matrix)
143
+
144
+ visual_prompt_images = [Image.fromarray(m) for m in filled_matrices]
145
+ return visual_prompt_images, selected_prompt_tokens, selected_prompt_indexes
146
+
147
+
148
+ def build_bboxes(masks: List[np.ndarray], tokenizer) -> Dict[str, tuple]:
149
+ height, width = masks[0].shape
150
+ bboxes: Dict[str, tuple] = {}
151
+ for idx, mask in enumerate(masks):
152
+ coords = np.argwhere(mask > 0)
153
+ if coords.size == 0:
154
+ continue
155
+ y_min, x_min = coords.min(axis=0)
156
+ y_max, x_max = coords.max(axis=0)
157
+ token_id = tokenizer.convert_tokens_to_ids(f"<|reserved_token_{idx}|>")
158
+ bboxes[str(token_id)] = (
159
+ x_min / width,
160
+ y_min / height,
161
+ x_max / width,
162
+ y_max / height,
163
+ )
164
+ return bboxes
165
+
166
+
167
+ def compute_aspect_ratio(image: Image.Image, processor, num_tiles: int) -> torch.Tensor:
168
+ min_tiles = getattr(processor, "min_sub_img", 1)
169
+ max_tiles = getattr(processor, "max_sub_img", 6)
170
+ if hasattr(processor, "image_size"):
171
+ image_size = processor.image_size[0] if isinstance(processor.image_size, tuple) else processor.image_size
172
+ else:
173
+ size = getattr(processor, "size", 512)
174
+ if isinstance(size, dict):
175
+ image_size = size.get("height", size.get("shortest_edge", 512))
176
+ else:
177
+ image_size = size
178
+ aspect_ratio = image.width / image.height
179
+ target_ratios = {
180
+ (i, j)
181
+ for n in range(min_tiles, max_tiles + 1)
182
+ for i in range(1, n + 1)
183
+ for j in range(1, n + 1)
184
+ if min_tiles <= i * j <= max_tiles
185
+ }
186
+ target_ratios = sorted(target_ratios, key=lambda x: x[0] * x[1])
187
+ grid_w, grid_h = find_closest_aspect_ratio(aspect_ratio, target_ratios, image.width, image.height, image_size)
188
+ return torch.tensor([[grid_w, grid_h]], dtype=torch.int64)
189
+
190
+
191
+ def build_prompt_text(tokenizer, num_image_token: int, num_tiles: int, questions: List[str], gen_len: int, num_masks: int) -> str:
192
+ img_ctx = "".join(["<IMG_CONTEXT>"] * (num_image_token * num_tiles))
193
+ parts = ["system\nYou are a helpful assistant.\n"]
194
+ parts.append("user\n")
195
+ parts.append(
196
+ f"<img>{img_ctx}</img>"
197
+ + "\n".join([f"<|reserved_token_{i}|>" for i in range(num_masks)])
198
+ + f"\n{questions[0]}\n"
199
+ )
200
+ parts.append("assistant\n")
201
+ mask_seq = "<|mdm_mask|>" * gen_len
202
+ parts.append("\n".join([f"<|Mask_Cap_{i}|>{mask_seq}" for i in range(num_masks)]))
203
+ return "".join(parts) + ""
204
+
205
+
206
+ def split_assistant_blocks(text: str, num_masks: int) -> List[str]:
207
+ blocks = text.split("assistant\n")
208
+ assistant_text = blocks[-1].split("")[0] if len(blocks) > 1 else text
209
+ captions = []
210
+ for i in range(num_masks):
211
+ start_tag = f"<|Mask_Cap_{i}|>"
212
+ next_tag = f"<|Mask_Cap_{i + 1}|>"
213
+ start_pos = assistant_text.find(start_tag)
214
+ if start_pos == -1:
215
+ captions.append("")
216
+ continue
217
+ content_start = start_pos + len(start_tag)
218
+ end_pos = assistant_text.find(next_tag, content_start) if i < num_masks - 1 else len(assistant_text)
219
+ if end_pos == -1:
220
+ end_pos = len(assistant_text)
221
+ captions.append(assistant_text[content_start:end_pos].strip())
222
+ return captions
223
+
224
+
225
+ # ---------------------------------------------------------------------------
226
+ # Helper utilities for the UI
227
+ # ---------------------------------------------------------------------------
228
+
229
+ def _to_binary_mask(mask_img: Image.Image, target_size: Tuple[int, int]) -> np.ndarray:
230
+ arr = np.array(mask_img.convert("L").resize(target_size, Image.NEAREST))
231
+ return (arr > 0).astype(np.uint8)
232
+
233
+
234
+ def make_overlay(pil_image: Image.Image, masks: List[np.ndarray], max_side: int = 768):
235
+ base = pil_image.convert("RGB")
236
+ w, h = base.size
237
+ scale = min(1.0, max_side / max(w, h))
238
+ if scale < 1.0:
239
+ new_size = (max(1, int(w * scale)), max(1, int(h * scale)))
240
+ base = base.resize(new_size, Image.BILINEAR)
241
+ annotations = []
242
+ for idx, mask in enumerate(masks):
243
+ m = mask.astype(np.uint8)
244
+ if scale < 1.0:
245
+ m = np.array(
246
+ Image.fromarray(m * 255).resize(base.size, Image.NEAREST)
247
+ ) > 0
248
+ m = m.astype(np.uint8)
249
+ annotations.append((m, f"Region {idx}"))
250
+ return (base, annotations)
251
+
252
+
253
+ def make_preset_thumbnail(image_path: str, mask_paths: List[str]) -> Image.Image:
254
+ img = Image.open(image_path).convert("RGB")
255
+ base = np.array(img).astype(np.float32)
256
+ for idx, mp in enumerate(mask_paths):
257
+ m = _to_binary_mask(Image.open(mp), img.size).astype(bool)
258
+ color = np.array(OVERLAY_COLORS[idx % len(OVERLAY_COLORS)], dtype=np.float32)
259
+ base[m] = 0.45 * base[m] + 0.55 * color
260
+ out = Image.fromarray(base.astype(np.uint8))
261
+ out.thumbnail((320, 320))
262
+ return out
263
+
264
+
265
+ # ---------------------------------------------------------------------------
266
+ # Decoding animation helpers
267
+ # ---------------------------------------------------------------------------
268
+
269
+ def decode_step_captions(step_tokens: torch.Tensor, num_masks: int) -> List[str]:
270
+ """Decode a single denoising step's token state into per-mask captions."""
271
+ ids = step_tokens[0].tolist()
272
+ pieces = []
273
+ for tid in ids:
274
+ if tid == MASK_ID:
275
+ pieces.append(MASK_PLACEHOLDER)
276
+ else:
277
+ pieces.append(TOKENIZER.decode([tid], skip_special_tokens=False))
278
+ raw = "".join(pieces)
279
+ captions = []
280
+ for i in range(num_masks):
281
+ start_tag = f"<|Mask_Cap_{i}|>"
282
+ next_tag = f"<|Mask_Cap_{i + 1}|>"
283
+ start_pos = raw.find(start_tag)
284
+ if start_pos == -1:
285
+ captions.append("")
286
+ continue
287
+ content_start = start_pos + len(start_tag)
288
+ end_pos = raw.find(next_tag, content_start) if i < num_masks - 1 else len(raw)
289
+ if end_pos == -1:
290
+ end_pos = len(raw)
291
+ text = raw[content_start:end_pos]
292
+ for tok in ("", "<|mdm_mask|>"):
293
+ text = text.replace(tok, "")
294
+ captions.append(text.strip())
295
+ return captions
296
+
297
+
298
+ def _render_caption_body(cap: str, prev_cap: str, color: tuple, highlight: bool = True) -> str:
299
+ rgb = f"rgb{color}"
300
+ out = []
301
+ prev_revealed = prev_cap.replace(MASK_PLACEHOLDER, "") if prev_cap else ""
302
+ seen_real = 0
303
+ for ch in cap:
304
+ if ch == MASK_PLACEHOLDER:
305
+ out.append(
306
+ f'<span class="tok-pending" style="background:{rgb};"></span>'
307
+ )
308
+ else:
309
+ seen_real += 1
310
+ is_new = highlight and seen_real > len(prev_revealed)
311
+ esc = html_lib.escape(ch)
312
+ if is_new:
313
+ out.append(f'<span class="tok-new" style="background:{rgb};">{esc}</span>')
314
+ else:
315
+ out.append(esc)
316
+ if not cap:
317
+ return '<span class="tok-empty">…</span>'
318
+ return "".join(out)
319
+
320
+
321
+ def render_caption_html(
322
+ captions: List[str],
323
+ prev_captions: List[str],
324
+ step_idx: int,
325
+ total_steps: int,
326
+ ) -> str:
327
+ last_step = max(total_steps - 1, 1)
328
+ is_final = step_idx >= total_steps - 1
329
+ pct = int(round(step_idx / last_step * 100))
330
+ css = """
331
+ <style>
332
+ .dec-wrap { font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif; }
333
+ .dec-head { display:flex; align-items:center; gap:12px; margin-bottom:14px; }
334
+ .dec-step { font-size:0.95em; font-weight:600; color:#475569; white-space:nowrap; }
335
+ .dec-progress { flex:1; height:6px; background:#e2e8f0; border-radius:3px; overflow:hidden; }
336
+ .dec-progress-fill { height:100%; background:linear-gradient(90deg,#6366f1,#a855f7); border-radius:3px; transition:width 0.25s ease; }
337
+ .cap-card { border:1px solid #e2e8f0; border-radius:12px; padding:14px 16px; margin-bottom:12px; background:#fff; box-shadow:0 1px 3px rgba(0,0,0,0.04); }
338
+ .cap-title { display:flex; align-items:center; gap:8px; font-weight:600; font-size:0.9em; margin-bottom:8px; color:#1e293b; }
339
+ .cap-dot { width:13px; height:13px; border-radius:50%; flex-shrink:0; }
340
+ .cap-body { font-size:0.95em; line-height:1.75; color:#0f172a; word-break:break-word; }
341
+ .tok-pending { display:inline-block; width:0.55em; height:0.55em; border-radius:50%; margin:0 1px; opacity:0.35; vertical-align:middle; animation:tokpulse 1.1s ease-in-out infinite; }
342
+ @keyframes tokpulse { 0%,100%{opacity:0.18;transform:scale(0.8);} 50%{opacity:0.6;transform:scale(1.05);} }
343
+ .tok-new { color:#fff; border-radius:4px; padding:0 2px; animation:tokreveal 0.45s ease-out; }
344
+ @keyframes tokreveal { from{opacity:0;transform:translateY(-3px) scale(0.9);} to{opacity:1;transform:none;} }
345
+ .tok-empty { color:#94a3b8; font-style:italic; }
346
+ </style>
347
+ """
348
+ parts = [css, '<div class="dec-wrap">']
349
+ parts.append(
350
+ f'<div class="dec-head"><span class="dec-step">Step {step_idx} / {last_step}</span>'
351
+ f'<div class="dec-progress"><div class="dec-progress-fill" style="width:{pct}%;"></div></div></div>'
352
+ )
353
+ for i, cap in enumerate(captions):
354
+ color = OVERLAY_COLORS[i % len(OVERLAY_COLORS)]
355
+ prev = prev_captions[i] if prev_captions and i < len(prev_captions) else ""
356
+ body = _render_caption_body(cap, prev, color, highlight=not is_final)
357
+ parts.append(
358
+ f'<div class="cap-card"><div class="cap-title">'
359
+ f'<span class="cap-dot" style="background:rgb{color};"></span>Region {i}</div>'
360
+ f'<div class="cap-body">{body}</div></div>'
361
+ )
362
+ parts.append("</div>")
363
+ return "".join(parts)
364
+
365
+
366
+ # ---------------------------------------------------------------------------
367
+ # GPU inference function (decorated with @spaces.GPU for ZeroGPU)
368
+ # ---------------------------------------------------------------------------
369
+
370
+ @spaces.GPU(duration=180)
371
+ def run_inference_gpu(
372
+ pil_image: Image.Image,
373
+ mask_images: List[Image.Image],
374
+ prompt: str,
375
+ gen_length: int,
376
+ steps: int,
377
+ temperature: float,
378
+ top_p: float,
379
+ ) -> List[List[str]]:
380
+ """Run the full PerceptionDLM pipeline and return per-step decoding history.
381
+
382
+ Each element is a list of per-mask caption strings for that denoising step.
383
+ All CUDA tensors are decoded to text inside the GPU worker so only plain
384
+ Python data crosses the pickle boundary.
385
+ """
386
+ prompt = prompt or DEFAULT_PROMPT
387
+ target_size = pil_image.size
388
+
389
+ masks_list = [_to_binary_mask(m, target_size) for m in mask_images]
390
+
391
+ sub_images = dynamic_preprocess(
392
+ pil_image,
393
+ min_num=PROCESSOR.min_sub_img,
394
+ max_num=PROCESSOR.max_sub_img,
395
+ image_size=PROCESSOR.image_size[0],
396
+ use_thumbnail=True,
397
+ )
398
+ pixel_values = PROCESSOR.image_processor.preprocess(
399
+ images=sub_images, return_tensors="pt"
400
+ )["pixel_values"].to("cuda").to(DTYPE)
401
+ aspect_ratio = compute_aspect_ratio(
402
+ pil_image, PROCESSOR, num_tiles=pixel_values.shape[0]
403
+ ).to("cuda")
404
+
405
+ sort_idx = sort_masks_by_area(masks_list)
406
+ masks_list = [masks_list[i] for i in sort_idx]
407
+
408
+ bboxes = build_bboxes(masks_list, TOKENIZER)
409
+ visual_prompt_images, prompt_tokens, _ = build_visual_prompt_matrices(
410
+ masks_list, prompt_numbers=MODEL.config.prompt_numbers
411
+ )
412
+
413
+ mask_values_list = []
414
+ for vp_img in visual_prompt_images:
415
+ vp_rgb = vp_img.convert("RGB")
416
+ sub_masks = dynamic_preprocess(
417
+ vp_rgb,
418
+ min_num=PROCESSOR.min_sub_img,
419
+ max_num=PROCESSOR.max_sub_img,
420
+ image_size=PROCESSOR.image_size[0],
421
+ use_thumbnail=True,
422
+ )
423
+ mv = PROCESSOR.image_processor.preprocess(
424
+ images=sub_masks, return_tensors="pt"
425
+ )["pixel_values"].to("cuda").to(DTYPE)
426
+ mask_values_list.append(mv)
427
+
428
+ questions = [prompt for _ in masks_list]
429
+ prompt_text = build_prompt_text(
430
+ tokenizer=TOKENIZER,
431
+ num_image_token=MODEL.config.num_image_token,
432
+ num_tiles=pixel_values.shape[0],
433
+ questions=questions,
434
+ gen_len=gen_length,
435
+ num_masks=len(masks_list),
436
+ )
437
+
438
+ model_inputs = TOKENIZER(prompt_text, return_tensors="pt")
439
+ input_ids = model_inputs["input_ids"].to("cuda")
440
+
441
+ _, all_steps = MODEL.generate_replace_noise(
442
+ pixel_values=pixel_values,
443
+ global_mask_values_list=mask_values_list,
444
+ aspect_ratios=aspect_ratio,
445
+ bboxes=[bboxes],
446
+ input_ids=input_ids,
447
+ steps=steps,
448
+ temperature=temperature,
449
+ top_p=top_p,
450
+ tokenizer=TOKENIZER,
451
+ prompt_tokens=prompt_tokens,
452
+ )
453
+
454
+ num_masks = len(masks_list)
455
+ history = [decode_step_captions(step_tok, num_masks) for step_tok in all_steps]
456
+ return history
457
+
458
+
459
+ # ---------------------------------------------------------------------------
460
+ # Preset examples
461
+ # ---------------------------------------------------------------------------
462
+ PRESET_DIR = os.path.dirname(os.path.abspath(__file__))
463
+ PRESETS: Dict[str, dict] = {}
464
+ demo_img = os.path.join(PRESET_DIR, "demo.jpg")
465
+ if os.path.exists(demo_img):
466
+ masks = sorted(
467
+ os.path.join(PRESET_DIR, f)
468
+ for f in os.listdir(PRESET_DIR)
469
+ if f.startswith("demo_mask_") and f.endswith(".jpg")
470
+ )
471
+ if masks:
472
+ PRESETS["demo.jpg with 3 masks"] = {"image": demo_img, "masks": masks}
473
+
474
+ PRESET_KEYS = list(PRESETS.keys())
475
+
476
+ # ---------------------------------------------------------------------------
477
+ # Build the Gradio interface
478
+ # ---------------------------------------------------------------------------
479
+
480
+ CUSTOM_CSS = """
481
+ #col-container { max-width: 1200px; margin: 0 auto; }
482
+ .dark .gradio-container { color: var(--body-text-color); }
483
+ .region-anno { overflow:hidden; }
484
+ .region-anno img, .region-anno canvas { max-width:100%; height:auto; object-fit:contain; }
485
+ """
486
+
487
+ color_map = {
488
+ f"Region {i}": "#%02x%02x%02x" % OVERLAY_COLORS[i % len(OVERLAY_COLORS)]
489
+ for i in range(len(OVERLAY_COLORS))
490
+ }
491
+
492
+ with gr.Blocks(
493
+ title="PerceptionDLM Region Captioning",
494
+ theme=gr.themes.Citrus(),
495
+ css=CUSTOM_CSS,
496
+ ) as demo:
497
+ gr.Markdown(
498
+ "# 🎯 PerceptionDLM Region Captioning\n"
499
+ "A diffusion multimodal LLM that captions any region of an image **in parallel**. "
500
+ "Upload an image and one or more binary masks, then run inference — "
501
+ "hover over a region to highlight it, and replay the diffusion decoding to watch each "
502
+ "caption emerge token by token.\n\n"
503
+ "Model: [MSALab/PerceptionDLM](https://huggingface.co/MSALab/PerceptionDLM) · "
504
+ "Paper: [arXiv:2606.19534](https://arxiv.org/abs/2606.19534) · "
505
+ "Code: [GitHub](https://github.com/MSALab-PKU/PerceptionDLM)"
506
+ )
507
+
508
+ with gr.Row(elem_id="col-container"):
509
+ with gr.Column(scale=1):
510
+ gr.Markdown("### Input")
511
+ if PRESETS:
512
+ preset_gallery = gr.Gallery(
513
+ value=[
514
+ (make_preset_thumbnail(c["image"], c["masks"]), name)
515
+ for name, c in PRESETS.items()
516
+ ],
517
+ columns=3,
518
+ height="auto",
519
+ object_fit="cover",
520
+ allow_preview=False,
521
+ label=None,
522
+ show_label=False,
523
+ )
524
+ gr.Markdown("*Click a thumbnail to load preset*")
525
+ image_in = gr.Image(type="pil", label="Image", image_mode="RGB")
526
+ mask_in = gr.File(
527
+ file_count="multiple",
528
+ file_types=["image"],
529
+ label="Mask images (binary, ≥1)",
530
+ )
531
+ prompt_in = gr.Textbox(value=DEFAULT_PROMPT, label="Prompt")
532
+ with gr.Accordion("Advanced settings", open=False):
533
+ with gr.Row():
534
+ gen_len_in = gr.Slider(8, 128, value=64, step=8, label="Gen length")
535
+ steps_in = gr.Slider(8, 128, value=32, step=8, label="Steps")
536
+ run_btn = gr.Button("Run inference", variant="primary")
537
+
538
+ with gr.Column(scale=1):
539
+ gr.Markdown("### Output")
540
+ overlay_out = gr.AnnotatedImage(
541
+ label="Regions (hover to highlight)",
542
+ color_map=color_map,
543
+ elem_classes=["region-anno"],
544
+ )
545
+ with gr.Row():
546
+ step_slider = gr.Slider(
547
+ 0, 1, value=0, step=1, label="Decoding step",
548
+ interactive=True, scale=4,
549
+ )
550
+ play_btn = gr.Button("▶ Play", variant="secondary", scale=1)
551
+ captions_out = gr.HTML()
552
+
553
+ # State
554
+ history_state = gr.State([])
555
+ num_masks_state = gr.State(0)
556
+
557
+ # ---- Preset loading via gallery click ----
558
+ def load_preset(evt: gr.SelectData):
559
+ name = PRESET_KEYS[evt.index]
560
+ case = PRESETS[name]
561
+ img = Image.open(case["image"]).convert("RGB")
562
+ return img, case["masks"]
563
+
564
+ if PRESETS:
565
+ preset_gallery.select(
566
+ load_preset, inputs=None, outputs=[image_in, mask_in]
567
+ )
568
+
569
+ # ---- Run inference ----
570
+ def _on_run(image, mask_files, prompt, gen_len, steps):
571
+ """Run PerceptionDLM inference and return overlay + decoding animation."""
572
+ if image is None:
573
+ raise gr.Error("Please provide an image.")
574
+ if not mask_files:
575
+ raise gr.Error("Please provide at least one mask image.")
576
+ mask_paths = [f if isinstance(f, str) else f.name for f in mask_files]
577
+ mask_images = [Image.open(p) for p in mask_paths]
578
+
579
+ history = run_inference_gpu(
580
+ image, mask_images, prompt, int(gen_len), int(steps),
581
+ temperature=0.0, top_p=1.0,
582
+ )
583
+
584
+ # Rebuild masks_list for overlay (same sorting as inside GPU fn)
585
+ target_size = image.size
586
+ masks_list = [_to_binary_mask(m, target_size) for m in mask_images]
587
+ sort_idx = sort_masks_by_area(masks_list)
588
+ masks_list = [masks_list[i] for i in sort_idx]
589
+ overlay = make_overlay(image, masks_list)
590
+
591
+ total = len(history)
592
+ last = total - 1
593
+ html = render_caption_html(
594
+ history[last], history[last - 1] if last > 0 else [], last, total
595
+ )
596
+ slider_update = gr.update(minimum=0, maximum=last, value=last, step=1)
597
+ return history, len(masks_list), overlay, slider_update, html
598
+
599
+ run_btn.click(
600
+ _on_run,
601
+ inputs=[image_in, mask_in, prompt_in, gen_len_in, steps_in],
602
+ outputs=[history_state, num_masks_state, overlay_out, step_slider, captions_out],
603
+ api_name="run_inference",
604
+ )
605
+
606
+ # ---- Step slider scrubbing ----
607
+ def _on_step(step_idx, history):
608
+ if not history:
609
+ return gr.update()
610
+ total = len(history)
611
+ i = int(step_idx)
612
+ i = max(0, min(i, total - 1))
613
+ prev = history[i - 1] if i > 0 else []
614
+ return render_caption_html(history[i], prev, i, total)
615
+
616
+ step_slider.change(_on_step, inputs=[step_slider, history_state], outputs=[captions_out])
617
+
618
+ # ---- Play animation ----
619
+ def _on_play(history):
620
+ if not history:
621
+ yield gr.update(), gr.update()
622
+ return
623
+ total = len(history)
624
+ for i in range(total):
625
+ prev = history[i - 1] if i > 0 else []
626
+ html = render_caption_html(history[i], prev, i, total)
627
+ yield gr.update(value=i), html
628
+ if i < total - 1:
629
+ time.sleep(0.25)
630
+
631
+ play_btn.click(_on_play, inputs=[history_state], outputs=[step_slider, captions_out])
632
+
633
+ # ---- Examples ----
634
+ example_entries = []
635
+ for name, case in PRESETS.items():
636
+ example_entries.append([case["image"]] + case["masks"] + [DEFAULT_PROMPT])
637
+
638
+ if example_entries:
639
+ gr.Examples(
640
+ examples=example_entries,
641
+ inputs=[image_in, mask_in, prompt_in],
642
+ fn=None, # examples fill inputs; user clicks Run
643
+ cache_examples=False,
644
+ run_on_click=True,
645
+ )
646
+
647
+
648
+ demo.queue()
649
+ if __name__ == "__main__":
650
+ demo.launch(mcp_server=True)
demo.jpg ADDED

Git LFS Details

  • SHA256: 3aeafa89d064257f4ba52f7e3b82f8a26b94fbc0e9a70445a37f0bc202f26555
  • Pointer size: 131 Bytes
  • Size of remote file: 113 kB
demo_mask_0.jpg ADDED
demo_mask_1.jpg ADDED
demo_mask_2.jpg ADDED
requirements.txt ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ transformers==4.51.3
2
+ accelerate
3
+ sentencepiece
4
+ einops
5
+ numpy
6
+ pillow
7
+ safetensors
8
+ torchvision