multimodalart HF Staff commited on
Commit
4607263
·
verified ·
1 Parent(s): 401b937

UI: Citrus theme, image-first gallery, randomize seed, β=1.2 default, cached examples

Browse files
.gitattributes CHANGED
@@ -33,3 +33,9 @@ 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
+ examples/chinese_vases/vase_03.jpg filter=lfs diff=lfs merge=lfs -text
37
+ examples/pink_elephant/bank_0000.png filter=lfs diff=lfs merge=lfs -text
38
+ examples/pink_elephant/bank_0001.png filter=lfs diff=lfs merge=lfs -text
39
+ examples/pink_elephant/bank_0002.png filter=lfs diff=lfs merge=lfs -text
40
+ examples/pink_elephant/bank_0003.png filter=lfs diff=lfs merge=lfs -text
41
+ examples/pink_elephant/bank_0004.png filter=lfs diff=lfs merge=lfs -text
app.py CHANGED
@@ -5,8 +5,8 @@ Code: https://github.com/pedrocurvo/follow-the-mean
5
 
6
  The core idea: in flow matching, the velocity field is governed by an endpoint
7
  mean. Shift that endpoint mean toward a reference set and the flow follows it.
8
- This Space exposes the training-free RMG variant: pick a reference (a prompt
9
- that will be sampled M times, or a folder of images), then guide a frozen
10
  FLUX.2-klein generator with the empirical reference endpoint mean.
11
  """
12
 
@@ -30,6 +30,8 @@ import retrieval_guidance_core as poc
30
  MODEL_ID = "black-forest-labs/FLUX.2-klein-4B"
31
  DEVICE = "cuda"
32
  DTYPE = torch.bfloat16
 
 
33
 
34
  # Load the pipeline at module scope. ZeroGPU intercepts CUDA at import time,
35
  # but the weights still need to be resident when the @spaces.GPU function runs.
@@ -135,12 +137,29 @@ def _make_grid(images: List[Image.Image], cell: int = 256) -> Image.Image:
135
  return grid
136
 
137
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
138
  @spaces.GPU(duration=300)
139
  def run_rmg(
140
  prompt: str,
141
  reference_mode: str,
142
  reference_prompt: str,
143
- reference_images: Optional[List[Image.Image]],
144
  reference_size: int,
145
  num_inference_steps: int,
146
  guidance_strength: float,
@@ -150,17 +169,22 @@ def run_rmg(
150
  topk: int,
151
  height: int,
152
  width: int,
 
153
  seed: int,
154
  progress=gr.Progress(track_tqdm=True),
155
  ):
156
  if not prompt or not prompt.strip():
157
  raise gr.Error("Please enter a prompt.")
158
 
159
- reference_mode = reference_mode.lower()
 
 
 
 
160
  if reference_mode.startswith("image"):
161
- if not reference_images:
162
- raise gr.Error("Upload at least one reference image, or switch to prompt mode.")
163
- ref_pils = [Image.open(p) if isinstance(p, str) else p for p in reference_images]
164
  ref_pils = _prepare_image_reference(ref_pils, width, height)
165
  else:
166
  rp = (reference_prompt or "").strip()
@@ -220,7 +244,7 @@ def run_rmg(
220
  ).images[0]
221
 
222
  reference_grid = _make_grid(ref_pils)
223
- return baseline, guided, reference_grid
224
 
225
 
226
  DESCRIPTION = """
@@ -233,49 +257,48 @@ goes. Change the reference set, and the flow changes.
233
 
234
  * Paper: [Follow the Mean: Reference-Guided Flow Matching](https://arxiv.org/abs/2605.10302)
235
  * Code: [github.com/pedrocurvo/follow-the-mean](https://github.com/pedrocurvo/follow-the-mean)
 
236
 
237
- **Reference mode**
238
- * *Prompt* — sample N reference images from the given reference prompt.
239
- * *Images* — upload your own reference photos.
240
 
241
- Each run produces the **baseline** (prompt only), the **RMG-guided** image
242
- (same prompt, same base noise, plus reference-mean correction), and a grid of
243
- the references that drove the guidance.
244
- """
 
245
 
246
  EXAMPLES = [
247
  [
248
- "an elephant in a jungle",
249
- "Prompt",
250
- "a pink elephant",
251
- None,
252
- 4, 28, 0.5, "quadratic-decay", 0.15, 0.95, 4, 1024, 1024, 123,
253
- ],
254
- [
255
- "a cat",
256
- "Prompt",
257
- "a van gogh painting",
258
- None,
259
- 4, 28, 0.5, "quadratic-decay", 0.15, 0.95, 4, 1024, 1024, 123,
260
  ],
261
  [
262
- "an animal in a savanna",
263
- "Prompt",
264
- "a giraffe",
265
- None,
266
- 4, 28, 0.6, "quadratic-decay", 0.15, 0.95, 4, 1024, 1024, 7,
267
- ],
268
- [
269
- "a house in a forest",
270
- "Prompt",
271
- "a pencil sketch",
272
- None,
273
- 4, 28, 0.5, "bell", 0.10, 0.90, 4, 1024, 1024, 42,
274
  ],
275
  ]
276
 
277
 
278
- with gr.Blocks(title="Follow the Mean — FLUX.2", theme=gr.themes.Soft()) as demo:
279
  gr.Markdown(DESCRIPTION)
280
 
281
  with gr.Row():
@@ -290,31 +313,35 @@ with gr.Blocks(title="Follow the Mean — FLUX.2", theme=gr.themes.Soft()) as de
290
  with gr.Group():
291
  reference_mode = gr.Radio(
292
  label="Reference mode",
293
- choices=["Prompt", "Images"],
294
- value="Prompt",
 
 
 
 
 
 
 
 
 
295
  )
296
  reference_prompt = gr.Textbox(
297
  label="Reference prompt (used in Prompt mode)",
298
  value="a pink elephant",
299
  placeholder="The attribute / object / style to inject.",
300
  lines=1,
301
- )
302
- reference_images = gr.File(
303
- label="Reference images (used in Images mode)",
304
- file_count="multiple",
305
- file_types=["image"],
306
- type="filepath",
307
  visible=False,
308
  )
309
  reference_size = gr.Slider(
310
  label="Reference set size (prompt mode)",
311
  minimum=1, maximum=12, step=1, value=4,
 
312
  )
313
 
314
- with gr.Accordion("Guidance", open=True):
315
  guidance_strength = gr.Slider(
316
  label="Guidance strength (β scale)",
317
- minimum=0.0, maximum=2.0, step=0.05, value=0.5,
318
  )
319
  beta_schedule = gr.Radio(
320
  label="β schedule",
@@ -340,7 +367,8 @@ with gr.Blocks(title="Follow the Mean — FLUX.2", theme=gr.themes.Soft()) as de
340
  with gr.Row():
341
  height = gr.Slider(label="Height", minimum=512, maximum=1280, step=64, value=1024)
342
  width = gr.Slider(label="Width", minimum=512, maximum=1280, step=64, value=1024)
343
- seed = gr.Slider(label="Seed", minimum=0, maximum=2_147_483_647, step=1, value=123)
 
344
 
345
  run_btn = gr.Button("Generate", variant="primary")
346
 
@@ -351,7 +379,7 @@ with gr.Blocks(title="Follow the Mean — FLUX.2", theme=gr.themes.Soft()) as de
351
  reference_out = gr.Image(label="Reference set", interactive=False)
352
 
353
  def _toggle_reference_inputs(mode):
354
- is_images = mode.lower().startswith("image")
355
  return (
356
  gr.update(visible=is_images), # reference_images
357
  gr.update(visible=not is_images), # reference_prompt
@@ -364,48 +392,34 @@ with gr.Blocks(title="Follow the Mean — FLUX.2", theme=gr.themes.Soft()) as de
364
  outputs=[reference_images, reference_prompt, reference_size],
365
  )
366
 
367
- run_btn.click(
368
- run_rmg,
369
- inputs=[
370
- prompt,
371
- reference_mode,
372
- reference_prompt,
373
- reference_images,
374
- reference_size,
375
- num_inference_steps,
376
- guidance_strength,
377
- beta_schedule,
378
- guidance_start_frac,
379
- guidance_end_frac,
380
- topk,
381
- height,
382
- width,
383
- seed,
384
- ],
385
- outputs=[baseline_out, guided_out, reference_out],
386
- )
387
 
388
  gr.Examples(
389
  examples=EXAMPLES,
390
- inputs=[
391
- prompt,
392
- reference_mode,
393
- reference_prompt,
394
- reference_images,
395
- reference_size,
396
- num_inference_steps,
397
- guidance_strength,
398
- beta_schedule,
399
- guidance_start_frac,
400
- guidance_end_frac,
401
- topk,
402
- height,
403
- width,
404
- seed,
405
- ],
406
- outputs=[baseline_out, guided_out, reference_out],
407
  fn=run_rmg,
408
- cache_examples=False,
 
409
  )
410
 
411
 
 
5
 
6
  The core idea: in flow matching, the velocity field is governed by an endpoint
7
  mean. Shift that endpoint mean toward a reference set and the flow follows it.
8
+ This Space exposes the training-free RMG variant: pick a reference (a folder
9
+ of images or a prompt that will be sampled M times), then guide a frozen
10
  FLUX.2-klein generator with the empirical reference endpoint mean.
11
  """
12
 
 
30
  MODEL_ID = "black-forest-labs/FLUX.2-klein-4B"
31
  DEVICE = "cuda"
32
  DTYPE = torch.bfloat16
33
+ MAX_SEED = 2_147_483_647
34
+ HERE = Path(__file__).resolve().parent
35
 
36
  # Load the pipeline at module scope. ZeroGPU intercepts CUDA at import time,
37
  # but the weights still need to be resident when the @spaces.GPU function runs.
 
137
  return grid
138
 
139
 
140
+ def _coerce_gallery_value(value) -> List[Image.Image]:
141
+ """Normalize a gr.Gallery value to a list of PIL images."""
142
+ if value is None:
143
+ return []
144
+ items = []
145
+ for entry in value:
146
+ if isinstance(entry, (list, tuple)):
147
+ entry = entry[0]
148
+ if isinstance(entry, dict):
149
+ entry = entry.get("image") or entry.get("name") or entry.get("path")
150
+ if isinstance(entry, Image.Image):
151
+ items.append(entry)
152
+ elif isinstance(entry, (str, Path)):
153
+ items.append(Image.open(entry))
154
+ return items
155
+
156
+
157
  @spaces.GPU(duration=300)
158
  def run_rmg(
159
  prompt: str,
160
  reference_mode: str,
161
  reference_prompt: str,
162
+ reference_images,
163
  reference_size: int,
164
  num_inference_steps: int,
165
  guidance_strength: float,
 
169
  topk: int,
170
  height: int,
171
  width: int,
172
+ randomize_seed: bool,
173
  seed: int,
174
  progress=gr.Progress(track_tqdm=True),
175
  ):
176
  if not prompt or not prompt.strip():
177
  raise gr.Error("Please enter a prompt.")
178
 
179
+ if randomize_seed:
180
+ seed = random.randint(0, MAX_SEED)
181
+ seed = int(seed)
182
+
183
+ reference_mode = (reference_mode or "").lower()
184
  if reference_mode.startswith("image"):
185
+ ref_pils = _coerce_gallery_value(reference_images)
186
+ if not ref_pils:
187
+ raise gr.Error("Add at least one reference image, or switch to prompt mode.")
188
  ref_pils = _prepare_image_reference(ref_pils, width, height)
189
  else:
190
  rp = (reference_prompt or "").strip()
 
244
  ).images[0]
245
 
246
  reference_grid = _make_grid(ref_pils)
247
+ return baseline, guided, reference_grid, seed
248
 
249
 
250
  DESCRIPTION = """
 
257
 
258
  * Paper: [Follow the Mean: Reference-Guided Flow Matching](https://arxiv.org/abs/2605.10302)
259
  * Code: [github.com/pedrocurvo/follow-the-mean](https://github.com/pedrocurvo/follow-the-mean)
260
+ """
261
 
 
 
 
262
 
263
+ PINK_ELEPHANT_IMAGES = sorted(str(p) for p in (HERE / "examples" / "pink_elephant").glob("*.png"))
264
+ CHINESE_VASE_IMAGES = sorted(
265
+ str(p) for p in (HERE / "examples" / "chinese_vases").glob("*.*")
266
+ )
267
+
268
 
269
  EXAMPLES = [
270
  [
271
+ "an elephant in a jungle", # prompt
272
+ "Images", # reference_mode
273
+ "", # reference_prompt (unused in image mode)
274
+ PINK_ELEPHANT_IMAGES, # reference_images (gallery)
275
+ 4, # reference_size (unused in image mode)
276
+ 28, # num_inference_steps
277
+ 1.2, # guidance_strength
278
+ "quadratic-decay", # beta_schedule
279
+ 0.15, 0.95, # start / end frac
280
+ 4, # topk
281
+ 1024, 1024, # height, width
282
+ False, 123, # randomize_seed, seed
283
  ],
284
  [
285
+ "a capybara",
286
+ "Images",
287
+ "",
288
+ CHINESE_VASE_IMAGES,
289
+ 4,
290
+ 28,
291
+ 1.25,
292
+ "quadratic-decay",
293
+ 0.15, 0.95,
294
+ 4,
295
+ 1024, 1024,
296
+ False, 42,
297
  ],
298
  ]
299
 
300
 
301
+ with gr.Blocks(title="Follow the Mean — FLUX.2", theme=gr.themes.Citrus()) as demo:
302
  gr.Markdown(DESCRIPTION)
303
 
304
  with gr.Row():
 
313
  with gr.Group():
314
  reference_mode = gr.Radio(
315
  label="Reference mode",
316
+ choices=["Images", "Prompt"],
317
+ value="Images",
318
+ )
319
+ reference_images = gr.Gallery(
320
+ label="Reference images",
321
+ columns=4,
322
+ height=240,
323
+ object_fit="cover",
324
+ type="filepath",
325
+ interactive=True,
326
+ visible=True,
327
  )
328
  reference_prompt = gr.Textbox(
329
  label="Reference prompt (used in Prompt mode)",
330
  value="a pink elephant",
331
  placeholder="The attribute / object / style to inject.",
332
  lines=1,
 
 
 
 
 
 
333
  visible=False,
334
  )
335
  reference_size = gr.Slider(
336
  label="Reference set size (prompt mode)",
337
  minimum=1, maximum=12, step=1, value=4,
338
+ visible=False,
339
  )
340
 
341
+ with gr.Accordion("Guidance", open=False):
342
  guidance_strength = gr.Slider(
343
  label="Guidance strength (β scale)",
344
+ minimum=0.0, maximum=2.0, step=0.05, value=1.2,
345
  )
346
  beta_schedule = gr.Radio(
347
  label="β schedule",
 
367
  with gr.Row():
368
  height = gr.Slider(label="Height", minimum=512, maximum=1280, step=64, value=1024)
369
  width = gr.Slider(label="Width", minimum=512, maximum=1280, step=64, value=1024)
370
+ randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
371
+ seed = gr.Slider(label="Seed", minimum=0, maximum=MAX_SEED, step=1, value=123)
372
 
373
  run_btn = gr.Button("Generate", variant="primary")
374
 
 
379
  reference_out = gr.Image(label="Reference set", interactive=False)
380
 
381
  def _toggle_reference_inputs(mode):
382
+ is_images = (mode or "").lower().startswith("image")
383
  return (
384
  gr.update(visible=is_images), # reference_images
385
  gr.update(visible=not is_images), # reference_prompt
 
392
  outputs=[reference_images, reference_prompt, reference_size],
393
  )
394
 
395
+ inputs_list = [
396
+ prompt,
397
+ reference_mode,
398
+ reference_prompt,
399
+ reference_images,
400
+ reference_size,
401
+ num_inference_steps,
402
+ guidance_strength,
403
+ beta_schedule,
404
+ guidance_start_frac,
405
+ guidance_end_frac,
406
+ topk,
407
+ height,
408
+ width,
409
+ randomize_seed,
410
+ seed,
411
+ ]
412
+ outputs_list = [baseline_out, guided_out, reference_out, seed]
413
+
414
+ run_btn.click(run_rmg, inputs=inputs_list, outputs=outputs_list)
415
 
416
  gr.Examples(
417
  examples=EXAMPLES,
418
+ inputs=inputs_list,
419
+ outputs=outputs_list,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
420
  fn=run_rmg,
421
+ cache_examples=True,
422
+ cache_mode="lazy",
423
  )
424
 
425
 
examples/chinese_vases/vase_01.jpeg ADDED
examples/chinese_vases/vase_02.webp ADDED
examples/chinese_vases/vase_03.jpg ADDED

Git LFS Details

  • SHA256: f4974c237a79975ad7af7cd1e4637d61b77e4d4070b4ccde28b4270e49503ce3
  • Pointer size: 132 Bytes
  • Size of remote file: 1.73 MB
examples/pink_elephant/bank_0000.png ADDED

Git LFS Details

  • SHA256: 069c7d4ddd7e4fd828cc6e453a2ba41b0ec769ec2f8caed056f670d2d750c309
  • Pointer size: 131 Bytes
  • Size of remote file: 406 kB
examples/pink_elephant/bank_0001.png ADDED

Git LFS Details

  • SHA256: 4460bf522b3d0d1aeed134beab05c843d0affde4816a1ba1dcaf8a844c75e12a
  • Pointer size: 131 Bytes
  • Size of remote file: 579 kB
examples/pink_elephant/bank_0002.png ADDED

Git LFS Details

  • SHA256: f9d551ed0b9c06e1e41748a2160a385bd0d4d15426a186f582393a3c325cf7ee
  • Pointer size: 131 Bytes
  • Size of remote file: 665 kB
examples/pink_elephant/bank_0003.png ADDED

Git LFS Details

  • SHA256: 87a2e4a45acce439ab8202cad00f1481b33e41500767697d21c8863ee2ce50c1
  • Pointer size: 131 Bytes
  • Size of remote file: 594 kB
examples/pink_elephant/bank_0004.png ADDED

Git LFS Details

  • SHA256: ce655c3dc971a161fbca14be3b602b9076abd79535e92dcc80b7fb60741078de
  • Pointer size: 131 Bytes
  • Size of remote file: 650 kB