MazCodes commited on
Commit
9571865
·
verified ·
1 Parent(s): df3e063

Upload folder using huggingface_hub

Browse files
.dockerignore CHANGED
@@ -4,7 +4,7 @@
4
  # docker-entrypoint.sh is COPY'd at build time and is the image's ENTRYPOINT.
5
 
6
  .git
7
-
8
  __pycache__/
9
  *.py[cod]
10
  *$py.class
 
4
  # docker-entrypoint.sh is COPY'd at build time and is the image's ENTRYPOINT.
5
 
6
  .git
7
+ docker-push.md
8
  __pycache__/
9
  *.py[cod]
10
  *$py.class
app/backend/app.py CHANGED
@@ -66,6 +66,38 @@ _components_initialised = False
66
  _init_error = None
67
 
68
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
69
  def _ensure_components():
70
  global config, audio_processor, generator, model_manager
71
  global _components_initialised, _init_error
@@ -339,6 +371,8 @@ def generate_audio():
339
  data.get('duration', 10.0), 'duration', min_value=1, max_value=60)
340
  cfg_scale = Validator.number(
341
  data.get('cfg_scale', 7.0), 'cfg_scale', min_value=0.1, max_value=20.0)
 
 
342
  seed = Validator.number(
343
  data.get('seed', -1), 'seed', min_value=-1, max_value=2**32 - 1, integer_only=True)
344
  batch_index = Validator.number(
@@ -349,8 +383,20 @@ def generate_audio():
349
  model_path = data.get('model_path')
350
  unwrapped_model_path = data.get('unwrapped_model_path')
351
 
 
 
 
 
 
 
 
 
 
 
352
  except ValidationError as e:
353
- return jsonify(APIResponse.validation_error({e.details['field']: [str(e)]})), 400
 
 
354
 
355
  logger.info(f"Audio generation request received")
356
  logger.debug(f"Request details: prompt='{prompt[:50]}...', duration={duration}s, model={model_name}")
@@ -395,6 +441,19 @@ def generate_audio():
395
  config_file, determined_model_path = determine_model_config(
396
  model_name, model_path, unwrapped_model_path)
397
  logger.info(f"Starting generation with config: {config_file}")
 
 
 
 
 
 
 
 
 
 
 
 
 
398
  try:
399
  if determined_model_path and determined_model_path.exists():
400
  output_path = generator.generate_audio(
@@ -404,9 +463,11 @@ def generate_audio():
404
  config_file=config_file,
405
  duration=duration,
406
  cfg_scale=cfg_scale,
 
407
  seed=seed,
408
  batch_index=batch_index,
409
- batch_total=batch_total
 
410
  )
411
  elif model_name in ['stable-audio-open-small', 'stable-audio-open-1.0']:
412
  model_file_mapping = {
@@ -427,9 +488,11 @@ def generate_audio():
427
  config_file=config_file,
428
  duration=duration,
429
  cfg_scale=cfg_scale,
 
430
  seed=seed,
431
  batch_index=batch_index,
432
- batch_total=batch_total
 
433
  )
434
  elif model_name and model_name != 'default':
435
  fine_tuned_path = config.get_path("models_fine_tuned") / model_name
@@ -438,17 +501,35 @@ def generate_audio():
438
 
439
  output_path = generator.generate_audio(
440
  prompt, fine_tuned_path, duration=duration,
441
- cfg_scale=cfg_scale, seed=seed,
442
- batch_index=batch_index, batch_total=batch_total)
 
443
  else:
444
  logger.debug("Using default model")
445
  output_path = generator.generate_audio(
446
- prompt, duration=duration, cfg_scale=cfg_scale, seed=seed,
447
- batch_index=batch_index, batch_total=batch_total)
 
448
 
449
  if not output_path.exists():
450
  raise GenerationError(prompt, model_name, "Generated audio file not found")
451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
452
  logger.info(f"Audio generation completed: {output_path.name} ({output_path.stat().st_size} bytes)")
453
  return send_file(
454
  str(output_path),
@@ -599,8 +680,10 @@ def get_models():
599
  unwrapped_models.sort(
600
  key=lambda x: x['created'], reverse=True)
601
 
602
- # Fine-tuned models reuse the base model's config for unwrapping.
603
- base_config_path = "models/config/model_config_small.json"
 
 
604
 
605
  models.append({
606
  'name': model_dir.name,
 
66
  _init_error = None
67
 
68
 
69
+ _LEGACY_FINETUNED_CONFIG_PATH = "models/config/model_config_small.json"
70
+
71
+
72
+ def _resolve_finetuned_config_path(model_dir: Path, config) -> str:
73
+ """Pick the architecture config for a fine-tuned model.
74
+
75
+ New runs drop a per-run model_config.json (and a training_metadata.json
76
+ breadcrumb) into the model folder. Older runs from before that change have
77
+ neither, so fall back to the small base config to preserve existing
78
+ behavior.
79
+ """
80
+ per_run_config = model_dir / "model_config.json"
81
+ if per_run_config.exists():
82
+ try:
83
+ return str(per_run_config.relative_to(config.project_root))
84
+ except ValueError:
85
+ return str(per_run_config)
86
+
87
+ metadata_path = model_dir / "training_metadata.json"
88
+ if metadata_path.exists():
89
+ try:
90
+ with open(metadata_path, 'r') as f:
91
+ metadata = json.load(f)
92
+ base_config = metadata.get("base_config_path")
93
+ if base_config and (config.project_root / base_config).exists():
94
+ return base_config
95
+ except (OSError, json.JSONDecodeError):
96
+ pass
97
+
98
+ return _LEGACY_FINETUNED_CONFIG_PATH
99
+
100
+
101
  def _ensure_components():
102
  global config, audio_processor, generator, model_manager
103
  global _components_initialised, _init_error
 
371
  data.get('duration', 10.0), 'duration', min_value=1, max_value=60)
372
  cfg_scale = Validator.number(
373
  data.get('cfg_scale', 7.0), 'cfg_scale', min_value=0.1, max_value=20.0)
374
+ steps = Validator.number(
375
+ data.get('steps', 250), 'steps', min_value=1, max_value=500, integer_only=True)
376
  seed = Validator.number(
377
  data.get('seed', -1), 'seed', min_value=-1, max_value=2**32 - 1, integer_only=True)
378
  batch_index = Validator.number(
 
383
  model_path = data.get('model_path')
384
  unwrapped_model_path = data.get('unwrapped_model_path')
385
 
386
+ align_bars_raw = data.get('align_bars')
387
+ align_bpm_raw = data.get('align_bpm')
388
+ align_bars = Validator.number(
389
+ align_bars_raw, 'align_bars', min_value=1, max_value=64,
390
+ integer_only=True) if align_bars_raw is not None else None
391
+ align_bpm = Validator.number(
392
+ align_bpm_raw, 'align_bpm', min_value=20, max_value=300
393
+ ) if align_bpm_raw is not None else None
394
+ do_align = align_bars is not None and align_bpm is not None
395
+
396
  except ValidationError as e:
397
+ field = e.details.get('field', 'unknown') if e.details else 'unknown'
398
+ logger.warning(f"/api/generate validation failed on '{field}': {e}")
399
+ return jsonify(APIResponse.validation_error({field: [str(e)]})), 400
400
 
401
  logger.info(f"Audio generation request received")
402
  logger.debug(f"Request details: prompt='{prompt[:50]}...', duration={duration}s, model={model_name}")
 
441
  config_file, determined_model_path = determine_model_config(
442
  model_name, model_path, unwrapped_model_path)
443
  logger.info(f"Starting generation with config: {config_file}")
444
+
445
+ # In bars mode we need a little extra audio so the post-processor can
446
+ # onset-trim and tempo-warp without running short of the requested length.
447
+ # The generator caps duration to model.sample_size internally, so this
448
+ # never overshoots the model's natural length.
449
+ ALIGN_HEADROOM_SECONDS = 1.5
450
+ if do_align:
451
+ duration = duration + ALIGN_HEADROOM_SECONDS
452
+ logger.debug(
453
+ f"Bars-mode alignment requested: bars={align_bars}, bpm={align_bpm}; "
454
+ f"requesting {duration:.2f}s with headroom"
455
+ )
456
+
457
  try:
458
  if determined_model_path and determined_model_path.exists():
459
  output_path = generator.generate_audio(
 
463
  config_file=config_file,
464
  duration=duration,
465
  cfg_scale=cfg_scale,
466
+ steps=steps,
467
  seed=seed,
468
  batch_index=batch_index,
469
+ batch_total=batch_total,
470
+ loop_mode=do_align,
471
  )
472
  elif model_name in ['stable-audio-open-small', 'stable-audio-open-1.0']:
473
  model_file_mapping = {
 
488
  config_file=config_file,
489
  duration=duration,
490
  cfg_scale=cfg_scale,
491
+ steps=steps,
492
  seed=seed,
493
  batch_index=batch_index,
494
+ batch_total=batch_total,
495
+ loop_mode=do_align,
496
  )
497
  elif model_name and model_name != 'default':
498
  fine_tuned_path = config.get_path("models_fine_tuned") / model_name
 
501
 
502
  output_path = generator.generate_audio(
503
  prompt, fine_tuned_path, duration=duration,
504
+ cfg_scale=cfg_scale, steps=steps, seed=seed,
505
+ batch_index=batch_index, batch_total=batch_total,
506
+ loop_mode=do_align)
507
  else:
508
  logger.debug("Using default model")
509
  output_path = generator.generate_audio(
510
+ prompt, duration=duration, cfg_scale=cfg_scale, steps=steps,
511
+ seed=seed, batch_index=batch_index, batch_total=batch_total,
512
+ loop_mode=do_align)
513
 
514
  if not output_path.exists():
515
  raise GenerationError(prompt, model_name, "Generated audio file not found")
516
 
517
+ if do_align:
518
+ try:
519
+ from app.core.generation.audio_post_process import align_to_grid
520
+ align_to_grid(
521
+ output_path,
522
+ target_bpm=float(align_bpm),
523
+ target_bars=int(align_bars),
524
+ )
525
+ logger.info(
526
+ f"Aligned to grid: bars={align_bars}, bpm={align_bpm}"
527
+ )
528
+ except Exception as exc:
529
+ # Never fail the request because alignment failed — the user
530
+ # would rather have the raw clip than an error toast.
531
+ logger.warning(f"Grid alignment skipped after error: {exc}")
532
+
533
  logger.info(f"Audio generation completed: {output_path.name} ({output_path.stat().st_size} bytes)")
534
  return send_file(
535
  str(output_path),
 
680
  unwrapped_models.sort(
681
  key=lambda x: x['created'], reverse=True)
682
 
683
+ # Resolve the architecture config for this fine-tuned model.
684
+ # Order: per-run copy in the model folder, then training_metadata
685
+ # breadcrumb, then legacy fallback to the small base config.
686
+ base_config_path = _resolve_finetuned_config_path(model_dir, config)
687
 
688
  models.append({
689
  'name': model_dir.name,
app/core/audio/link_sync.py CHANGED
@@ -61,11 +61,32 @@ class LinkBridge:
61
  if ctor is not None:
62
  try:
63
  self._link = ctor(self._last_bpm)
 
 
 
 
64
  return
65
  except Exception as exc:
66
  logger.warning(f"Link ctor {ctor_name} failed: {exc}")
67
  logger.warning("Loaded Link module but no known constructor was found")
68
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
69
  def _set_enabled_on_link(self, value: bool) -> None:
70
  if self._link is None:
71
  return
@@ -114,17 +135,43 @@ class LinkBridge:
114
  "enabled": self._enabled,
115
  "bpm": self._last_bpm,
116
  "num_peers": 0,
 
 
 
117
  }
118
  if not self._enabled or self._link is None:
119
  return state
120
  try:
121
  session = self._capture_session()
122
- if session is not None and hasattr(session, "tempo"):
 
 
123
  state["bpm"] = float(session.tempo())
124
  self._last_bpm = state["bpm"]
125
  num_peers = getattr(self._link, "numPeers", None)
126
  if callable(num_peers):
127
  state["num_peers"] = int(num_peers())
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
128
  except Exception as exc:
129
  logger.warning(f"Link state read failed: {exc}")
130
  return state
 
61
  if ctor is not None:
62
  try:
63
  self._link = ctor(self._last_bpm)
64
+ # Tempo sync works without this; transport (play/stop) sync
65
+ # is opt-in per peer. Live exposes the same toggle as
66
+ # "Start Stop Sync" in its Link/Tempo/MIDI prefs.
67
+ self._set_start_stop_sync(True)
68
  return
69
  except Exception as exc:
70
  logger.warning(f"Link ctor {ctor_name} failed: {exc}")
71
  logger.warning("Loaded Link module but no known constructor was found")
72
 
73
+ def _set_start_stop_sync(self, value: bool) -> None:
74
+ if self._link is None:
75
+ return
76
+ try:
77
+ self._link.startStopSyncEnabled = value
78
+ return
79
+ except AttributeError:
80
+ pass
81
+ for setter_name in ("setStartStopSyncEnabled", "set_start_stop_sync_enabled"):
82
+ setter = getattr(self._link, setter_name, None)
83
+ if setter:
84
+ try:
85
+ setter(value)
86
+ return
87
+ except Exception as exc:
88
+ logger.warning(f"Link {setter_name} failed: {exc}")
89
+
90
  def _set_enabled_on_link(self, value: bool) -> None:
91
  if self._link is None:
92
  return
 
135
  "enabled": self._enabled,
136
  "bpm": self._last_bpm,
137
  "num_peers": 0,
138
+ "is_playing": False,
139
+ "beat": 0.0,
140
+ "time_micros": 0,
141
  }
142
  if not self._enabled or self._link is None:
143
  return state
144
  try:
145
  session = self._capture_session()
146
+ if session is None:
147
+ return state
148
+ if hasattr(session, "tempo"):
149
  state["bpm"] = float(session.tempo())
150
  self._last_bpm = state["bpm"]
151
  num_peers = getattr(self._link, "numPeers", None)
152
  if callable(num_peers):
153
  state["num_peers"] = int(num_peers())
154
+
155
+ for name in ("isPlaying", "is_playing"):
156
+ fn = getattr(session, name, None)
157
+ if callable(fn):
158
+ state["is_playing"] = bool(fn())
159
+ break
160
+
161
+ # Sample beat + host time together so the client can extrapolate
162
+ # forward by (now - capturedAt) when scheduling launches.
163
+ clock_fn = getattr(self._link, "clock", None)
164
+ if callable(clock_fn):
165
+ micros = int(clock_fn().micros())
166
+ state["time_micros"] = micros
167
+ # Quantum value here only affects phase wrapping for peers
168
+ # that just joined; for an existing session the absolute
169
+ # beat is stable. 4 = one bar in 4/4, a sane default.
170
+ for name in ("beatAtTime", "beat_at_time"):
171
+ fn = getattr(session, name, None)
172
+ if callable(fn):
173
+ state["beat"] = float(fn(micros, 4.0))
174
+ break
175
  except Exception as exc:
176
  logger.warning(f"Link state read failed: {exc}")
177
  return state
app/core/generation/audio_generator.py CHANGED
@@ -25,7 +25,7 @@ def _slugify_prompt(text: str, max_len: int = 40) -> str:
25
  sys.path.append(
26
  str(Path(__file__).parent.parent.parent.parent / "stable-audio-tools"))
27
 
28
- # Third-party package noise (clip/pkg_resources) is non-actionable for runtime.
29
  warnings.filterwarnings(
30
  "ignore",
31
  message=r"pkg_resources is deprecated as an API.*",
@@ -46,6 +46,11 @@ class AudioGenerator:
46
  self.device = "cuda" if torch.cuda.is_available() else "cpu"
47
  self.current_model_name = None
48
  self.current_model_path = None
 
 
 
 
 
49
  self._stop_event = threading.Event()
50
  logger.info(f"Using device: {self.device}")
51
 
@@ -67,7 +72,8 @@ class AudioGenerator:
67
  config_file = "model_config_small.json"
68
  else:
69
  config_file = "model_config.json"
70
-
 
71
  config_path = Path(__file__).parent.parent.parent.parent / "models" / "config" / config_file
72
  logger.info(f"Using config file: {config_path}")
73
 
@@ -136,6 +142,7 @@ class AudioGenerator:
136
  from stable_audio_tools.models.utils import load_ckpt_state_dict
137
  if config_file is None:
138
  config_file = "model_config_small.json"
 
139
 
140
  config_path = Path(__file__).parent.parent.parent.parent / \
141
  "models" / "config" / config_file
@@ -176,13 +183,25 @@ class AudioGenerator:
176
  seed: int = -1,
177
  output_path: Optional[Path] = None,
178
  batch_index: int = 1,
179
- batch_total: int = 1
 
180
  ) -> Path:
181
  print(f"\nAUDIO GENERATOR: generate_audio called")
182
  print(f" - Prompt: '{prompt}'")
183
  print(f" - Duration: {duration}s")
184
 
185
- if self.model is None or model_path is not None or unwrapped_model_path is not None:
 
 
 
 
 
 
 
 
 
 
 
186
  print(f"AUDIO GENERATOR: Loading new model")
187
 
188
  if unwrapped_model_path:
@@ -209,8 +228,8 @@ class AudioGenerator:
209
  print(f"AUDIO GENERATOR: Loading default local small base model")
210
  if not self.load_local_base_model("stable-audio-open-small"):
211
  raise ValueError("Failed to load default local base model")
212
- else:
213
- print(f"AUDIO GENERATOR: Using existing model")
214
 
215
  print(f"AUDIO GENERATOR: Model loaded successfully")
216
 
@@ -222,8 +241,26 @@ class AudioGenerator:
222
  raise GenerationStopped("Stop requested mid-diffusion")
223
 
224
  try:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
225
  print(f"Generating audio for prompt: '{prompt}'")
226
- print(f"Duration: {duration}s, CFG scale: {cfg_scale}, Steps: {steps}")
 
 
 
 
227
  requested_sample_size = int(duration * self.model.sample_rate)
228
  max_sample_size = None
229
  try:
@@ -269,10 +306,23 @@ class AudioGenerator:
269
  seed = np.random.randint(0, 2**32 - 1, dtype=np.int64)
270
 
271
  print(f"Using seed: {seed}")
 
 
 
 
 
 
 
 
 
 
 
 
 
272
  conditioning = [{
273
  "prompt": prompt,
274
  "seconds_start": 0,
275
- "seconds_total": int(duration)
276
  }]
277
 
278
  device = next(self.model.parameters()).device
@@ -295,17 +345,16 @@ class AudioGenerator:
295
 
296
  audio = generate_diffusion_cond(
297
  model=self.model,
298
- steps=steps,
299
- cfg_scale=cfg_scale,
300
  conditioning=conditioning,
301
  batch_size=1,
302
  sample_size=requested_sample_size,
303
  seed=seed,
304
  device=str(device),
305
- sigma_min=0.03,
306
- sigma_max=1000,
307
- sampler_type="dpmpp-3m-sde",
308
- callback=_stop_callback
309
  )
310
 
311
  print(f"Generation complete, audio shape: {audio.shape}")
 
25
  sys.path.append(
26
  str(Path(__file__).parent.parent.parent.parent / "stable-audio-tools"))
27
 
28
+
29
  warnings.filterwarnings(
30
  "ignore",
31
  message=r"pkg_resources is deprecated as an API.*",
 
46
  self.device = "cuda" if torch.cuda.is_available() else "cpu"
47
  self.current_model_name = None
48
  self.current_model_path = None
49
+ # Caller-facing key for the currently loaded weights. Used to short-
50
+ # circuit reloads when generate_audio is invoked repeatedly with the
51
+ # same model — without this, every Generate click reloads from disk.
52
+ self.current_model_key = None
53
+ self.is_distilled_small = False
54
  self._stop_event = threading.Event()
55
  logger.info(f"Using device: {self.device}")
56
 
 
72
  config_file = "model_config_small.json"
73
  else:
74
  config_file = "model_config.json"
75
+ self.is_distilled_small = "small" in model_name.lower()
76
+
77
  config_path = Path(__file__).parent.parent.parent.parent / "models" / "config" / config_file
78
  logger.info(f"Using config file: {config_path}")
79
 
 
142
  from stable_audio_tools.models.utils import load_ckpt_state_dict
143
  if config_file is None:
144
  config_file = "model_config_small.json"
145
+ self.is_distilled_small = "small" in config_file.lower()
146
 
147
  config_path = Path(__file__).parent.parent.parent.parent / \
148
  "models" / "config" / config_file
 
183
  seed: int = -1,
184
  output_path: Optional[Path] = None,
185
  batch_index: int = 1,
186
+ batch_total: int = 1,
187
+ loop_mode: bool = False,
188
  ) -> Path:
189
  print(f"\nAUDIO GENERATOR: generate_audio called")
190
  print(f" - Prompt: '{prompt}'")
191
  print(f" - Duration: {duration}s")
192
 
193
+ # Build a cache key for the requested model so we can reuse weights
194
+ # across consecutive Generate clicks on the same model.
195
+ if unwrapped_model_path:
196
+ target_key = ('unwrapped', str(unwrapped_model_path))
197
+ elif model_path:
198
+ target_key = ('path', str(model_path))
199
+ else:
200
+ target_key = ('default', 'stable-audio-open-small')
201
+
202
+ if self.model is not None and self.current_model_key == target_key:
203
+ print(f"AUDIO GENERATOR: Reusing already-loaded model")
204
+ else:
205
  print(f"AUDIO GENERATOR: Loading new model")
206
 
207
  if unwrapped_model_path:
 
228
  print(f"AUDIO GENERATOR: Loading default local small base model")
229
  if not self.load_local_base_model("stable-audio-open-small"):
230
  raise ValueError("Failed to load default local base model")
231
+
232
+ self.current_model_key = target_key
233
 
234
  print(f"AUDIO GENERATOR: Model loaded successfully")
235
 
 
241
  raise GenerationStopped("Stop requested mid-diffusion")
242
 
243
  try:
244
+ # Stable Audio Open Small is an adversarially-distilled checkpoint
245
+ # that requires the pingpong sampler at 8 steps with CFG 1.0.
246
+ # Running the dpmpp-3m-sde recipe on it produces noise.
247
+ if self.is_distilled_small:
248
+ effective_sampler = "pingpong"
249
+ effective_steps = 8
250
+ effective_cfg = 1.0
251
+ sigma_kwargs = {}
252
+ else:
253
+ effective_sampler = "dpmpp-3m-sde"
254
+ effective_steps = steps
255
+ effective_cfg = cfg_scale
256
+ sigma_kwargs = {"sigma_min": 0.03, "sigma_max": 1000}
257
+
258
  print(f"Generating audio for prompt: '{prompt}'")
259
+ print(
260
+ f"Duration: {duration}s, CFG scale: {effective_cfg}, "
261
+ f"Steps: {effective_steps}, Sampler: {effective_sampler}"
262
+ + (" (distilled small overrides applied)" if self.is_distilled_small else "")
263
+ )
264
  requested_sample_size = int(duration * self.model.sample_rate)
265
  max_sample_size = None
266
  try:
 
306
  seed = np.random.randint(0, 2**32 - 1, dtype=np.int64)
307
 
308
  print(f"Using seed: {seed}")
309
+
310
+ # In loop_mode (bars mode in the performance panel), tell the model
311
+ # the song is much longer than what we render. SAO 1.0 was trained
312
+ # to fade out as it approaches `seconds_total`, so matching it to
313
+ # the requested duration bakes a song-ending fade into the clip.
314
+ # We still get back exactly `requested_sample_size` samples — the
315
+ # model just thinks they're the opening of a longer piece.
316
+ if loop_mode and max_sample_size:
317
+ song_seconds = max(int(duration),
318
+ int(max_sample_size / self.model.sample_rate))
319
+ else:
320
+ song_seconds = int(duration)
321
+
322
  conditioning = [{
323
  "prompt": prompt,
324
  "seconds_start": 0,
325
+ "seconds_total": song_seconds,
326
  }]
327
 
328
  device = next(self.model.parameters()).device
 
345
 
346
  audio = generate_diffusion_cond(
347
  model=self.model,
348
+ steps=effective_steps,
349
+ cfg_scale=effective_cfg,
350
  conditioning=conditioning,
351
  batch_size=1,
352
  sample_size=requested_sample_size,
353
  seed=seed,
354
  device=str(device),
355
+ sampler_type=effective_sampler,
356
+ callback=_stop_callback,
357
+ **sigma_kwargs,
 
358
  )
359
 
360
  print(f"Generation complete, audio shape: {audio.shape}")
app/core/generation/audio_post_process.py ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Beat-align and tempo-conform a generated WAV to a target BPM and bar count.
2
+ """
3
+
4
+ from __future__ import annotations
5
+
6
+ import logging
7
+ from pathlib import Path
8
+ from typing import Optional, Tuple
9
+
10
+ import librosa
11
+ import numpy as np
12
+ import soundfile as sf
13
+
14
+ logger = logging.getLogger(__name__)
15
+
16
+
17
+ def align_to_grid(
18
+ input_path: Path,
19
+ target_bpm: float,
20
+ target_bars: int,
21
+ beats_per_bar: int = 4,
22
+ ) -> Path:
23
+ audio, sr = sf.read(str(input_path), always_2d=True)
24
+ audio = audio.astype(np.float32, copy=False)
25
+ target_samples = int(round(target_bars * beats_per_bar * 60.0 / target_bpm * sr))
26
+
27
+ mono = audio.mean(axis=1) if audio.shape[1] > 1 else audio[:, 0]
28
+
29
+ detected_bpm, first_beat = _detect_grid_anchor(mono, sr)
30
+
31
+ head_offset = 0
32
+ if first_beat is not None and 0 < first_beat < sr * 1.5:
33
+ head_offset = first_beat
34
+ logger.info(f"align_to_grid: trimmed {head_offset / sr * 1000:.1f} ms to first beat")
35
+ elif first_beat is None:
36
+ head_offset = _detect_first_onset_sample(mono, sr)
37
+ if head_offset > 0:
38
+ logger.info(f"align_to_grid: trimmed {head_offset / sr * 1000:.1f} ms (onset fallback)")
39
+
40
+ if head_offset > 0:
41
+ audio = audio[head_offset:]
42
+ mono = mono[head_offset:]
43
+
44
+ if detected_bpm is not None:
45
+ rate = target_bpm / detected_bpm
46
+ if 0.7 <= rate <= 1.4:
47
+ audio = _time_stretch_multichannel(audio, rate)
48
+ logger.info(
49
+ f"align_to_grid: detected {detected_bpm:.2f} BPM, "
50
+ f"stretched by {rate:.4f} to match target {target_bpm:.2f} BPM"
51
+ )
52
+ else:
53
+ logger.info(
54
+ f"align_to_grid: detected {detected_bpm:.2f} BPM out of safe stretch "
55
+ f"range vs target {target_bpm:.2f}; skipping warp"
56
+ )
57
+ else:
58
+ logger.info("align_to_grid: no usable tempo detected; skipping warp")
59
+
60
+ if audio.shape[0] > target_samples:
61
+ audio = audio[:target_samples]
62
+ elif audio.shape[0] < target_samples:
63
+ pad = np.zeros((target_samples - audio.shape[0], audio.shape[1]), dtype=audio.dtype)
64
+ audio = np.concatenate([audio, pad], axis=0)
65
+
66
+ sf.write(str(input_path), audio, sr, subtype="PCM_16")
67
+ return input_path
68
+
69
+
70
+ def _detect_first_onset_sample(mono: np.ndarray, sr: int) -> int:
71
+ """Return the sample index of the first detected onset, or 0 if none found."""
72
+ try:
73
+ onsets = librosa.onset.onset_detect(
74
+ y=mono, sr=sr, units="samples", backtrack=True
75
+ )
76
+ except Exception as exc:
77
+ logger.warning(f"onset detection failed: {exc}")
78
+ return 0
79
+ if onsets is None or len(onsets) == 0:
80
+ return 0
81
+ first = int(onsets[0])
82
+ if first > sr * 1.0:
83
+ return 0
84
+ return first
85
+
86
+
87
+ def _detect_grid_anchor(mono: np.ndarray, sr: int) -> Tuple[Optional[float], Optional[int]]:
88
+ try:
89
+ tempo, beats = librosa.beat.beat_track(y=mono, sr=sr, units="samples")
90
+ except Exception as exc:
91
+ logger.warning(f"beat tracking failed: {exc}")
92
+ return None, None
93
+ if beats is None or len(beats) < 4:
94
+ return None, None
95
+ bpm = float(np.atleast_1d(tempo).flatten()[0])
96
+ if not (40.0 <= bpm <= 240.0):
97
+ return None, None
98
+ return bpm, int(beats[0])
99
+
100
+
101
+ def _time_stretch_multichannel(audio: np.ndarray, rate: float) -> np.ndarray:
102
+ """Phase-vocoder time stretch, applied per channel and re-stacked."""
103
+ stretched = librosa.effects.time_stretch(audio.T, rate=rate)
104
+ return np.ascontiguousarray(stretched.T)
app/core/training/fine_tuner.py CHANGED
@@ -307,30 +307,59 @@ class FineTuner:
307
  learning_rate = self.config.get("learningRate", 1e-4)
308
  print(f"LEARNING RATE: {learning_rate}")
309
 
 
 
 
 
 
 
 
 
 
310
  import json
311
- config_path = Path(model_config)
312
- if config_path.exists():
313
- with open(config_path, 'r') as f:
 
314
  config_data = json.load(f)
315
 
316
- if 'training' in config_data and 'optimizer_configs' in config_data['training']:
317
- if 'diffusion' in config_data['training']['optimizer_configs']:
318
- if 'optimizer' in config_data['training']['optimizer_configs']['diffusion']:
319
- if 'config' in config_data['training']['optimizer_configs']['diffusion']['optimizer']:
320
- old_lr = config_data['training']['optimizer_configs']['diffusion']['optimizer']['config']['lr']
321
- config_data['training']['optimizer_configs']['diffusion']['optimizer']['config']['lr'] = learning_rate
322
- print(f"Updated learning rate from {old_lr} to {learning_rate} in model config")
323
-
324
- with open(config_path, 'w') as f:
 
 
325
  json.dump(config_data, f, indent=4)
326
- print(f"Updated model config saved to: {config_path}")
 
327
  else:
328
- print(f"WARNING: Model config file not found: {config_path}")
329
 
330
- config = get_config()
331
- dataset_config = config.get_dataset_config_path()
332
- save_dir = str(config.get_path("models_fine_tuned") / model_name)
333
- os.makedirs(save_dir, exist_ok=True)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
334
 
335
  batch_size = self.config.get("batchSize", 4)
336
  accum_batches = 1
 
307
  learning_rate = self.config.get("learningRate", 1e-4)
308
  print(f"LEARNING RATE: {learning_rate}")
309
 
310
+ config = get_config()
311
+ dataset_config = config.get_dataset_config_path()
312
+ save_dir = str(config.get_path("models_fine_tuned") / model_name)
313
+ os.makedirs(save_dir, exist_ok=True)
314
+
315
+ # Write a per-run copy of the model config into save_dir with the LR
316
+ # override applied. This avoids mutating the shared base config on
317
+ # disk and makes the fine-tuned model folder self-describing for
318
+ # later loading/unwrapping.
319
  import json
320
+ base_config_path = Path(model_config)
321
+ run_config_path = Path(save_dir) / "model_config.json"
322
+ if base_config_path.exists():
323
+ with open(base_config_path, 'r') as f:
324
  config_data = json.load(f)
325
 
326
+ try:
327
+ optimizer_config = (
328
+ config_data['training']['optimizer_configs']['diffusion']['optimizer']['config']
329
+ )
330
+ old_lr = optimizer_config.get('lr')
331
+ optimizer_config['lr'] = learning_rate
332
+ print(f"Updated learning rate from {old_lr} to {learning_rate} in run config")
333
+ except (KeyError, TypeError):
334
+ print(f"WARNING: Could not locate optimizer.lr in base config; LR override skipped")
335
+
336
+ with open(run_config_path, 'w') as f:
337
  json.dump(config_data, f, indent=4)
338
+ print(f"Per-run model config written to: {run_config_path}")
339
+ model_config = str(run_config_path)
340
  else:
341
+ print(f"WARNING: Base model config file not found: {base_config_path}")
342
 
343
+ # Drop a metadata breadcrumb so the read side (app.py /api/models)
344
+ # knows which base architecture this fine-tune is paired with,
345
+ # instead of guessing.
346
+ metadata_path = Path(save_dir) / "training_metadata.json"
347
+ try:
348
+ base_config_rel = str(base_config_path.relative_to(project_root))
349
+ except ValueError:
350
+ base_config_rel = str(base_config_path)
351
+ try:
352
+ pretrained_ckpt_rel = str(Path(pretrained_ckpt).relative_to(project_root))
353
+ except ValueError:
354
+ pretrained_ckpt_rel = pretrained_ckpt
355
+ with open(metadata_path, 'w') as f:
356
+ json.dump({
357
+ "base_model": base_model,
358
+ "base_config_path": base_config_rel,
359
+ "pretrained_ckpt_path": pretrained_ckpt_rel,
360
+ "learning_rate": learning_rate,
361
+ }, f, indent=4)
362
+ print(f"Training metadata written to: {metadata_path}")
363
 
364
  batch_size = self.config.get("batchSize", 4)
365
  accum_batches = 1
app/frontend/package.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "fragmenta-desktop",
3
- "version": "0.0.2",
4
  "description": "Fragmenta Desktop",
5
  "type": "module",
6
  "scripts": {
 
1
  {
2
  "name": "fragmenta-desktop",
3
+ "version": "0.1.0",
4
  "description": "Fragmenta Desktop",
5
  "type": "module",
6
  "scripts": {
app/frontend/public/fragmenta.ico ADDED
app/frontend/src/App.js CHANGED
@@ -59,6 +59,7 @@ import ModelUnwrapButton from './components/ModelUnwrapButton';
59
  import CheckpointManager from './components/CheckpointManager';
60
  import GeneratedFragmentsWindow from './components/GeneratedFragmentsWindow';
61
  import WelcomePage from './components/WelcomePage';
 
62
  import { formatDuration } from './utils/format';
63
  import theme, { appStyles, lightTheme } from './theme';
64
 
@@ -137,6 +138,7 @@ function App() {
137
  const [generatedFragments, setGeneratedFragments] = useState([]);
138
  const [currentFilename, setCurrentFilename] = useState('');
139
  const [cfgScale, setCfgScale] = useState(7.0);
 
140
  const [batchCount, setBatchCount] = useState(1);
141
  const [randomSeed, setRandomSeed] = useState(true);
142
  const [seedValue, setSeedValue] = useState('');
@@ -206,6 +208,10 @@ function App() {
206
  const [showStartFreshDialog, setShowStartFreshDialog] = useState(false);
207
  const [isStartingFresh, setIsStartingFresh] = useState(false);
208
  const [uploadKey, setUploadKey] = useState(0);
 
 
 
 
209
  const [isFreeingGPU, setIsFreeingGPU] = useState(false);
210
  const [showFreeGPUDialog, setShowFreeGPUDialog] = useState(false);
211
  const [modelWarning, setModelWarning] = useState({
@@ -253,6 +259,16 @@ function App() {
253
  return 10;
254
  };
255
 
 
 
 
 
 
 
 
 
 
 
256
  useEffect(() => {
257
  const maxDuration = getMaxDuration();
258
  if (generationDuration > maxDuration) {
@@ -523,7 +539,8 @@ function App() {
523
  const baseRequestData = {
524
  prompt: generationPrompt,
525
  duration: generationDuration,
526
- cfg_scale: cfgScale
 
527
  };
528
 
529
  const baseModel = baseModels.find(m => m.name === selectedModel);
@@ -631,6 +648,7 @@ function App() {
631
  prompt: generationPrompt,
632
  duration: generationDuration,
633
  cfgScale,
 
634
  seed: seedForRun,
635
  batchIndex,
636
  batchTotal: totalRuns,
@@ -682,9 +700,6 @@ function App() {
682
 
683
  const stopGeneration = () => {
684
  stopGenerationRef.current = true;
685
- // Tell the backend first so the in-flight diffusion loop bails out at
686
- // the next step, then abort the HTTP request on our side. Fire-and-
687
- // forget; we deliberately do not pass the abort signal here.
688
  api.post('/api/stop-generation').catch(() => {});
689
  if (generationAbortRef.current) {
690
  try { generationAbortRef.current.abort(); } catch (_) {}
@@ -709,6 +724,12 @@ function App() {
709
  setGenerationPrompt('');
710
  setUploadKey(prev => prev + 1);
711
 
 
 
 
 
 
 
712
  setProcessingStatus(response.data.message);
713
 
714
  fetchSystemStatus();
@@ -1080,10 +1101,10 @@ function App() {
1080
  <Paper sx={{ p: 2 }} variant="outlined">
1081
  <Box sx={{ display: 'flex', alignItems: 'center', gap: 1, mb: 1.5 }}>
1082
  <UploadIcon size={20} />
1083
- <Typography variant="h6">Manual Annotate</Typography>
1084
  </Box>
1085
  <Typography variant="body2" color="textSecondary" sx={{ mb: 2 }}>
1086
- Upload audio files one by one and write each prompt yourself.
1087
  Use this when you want full control over every annotation.
1088
  </Typography>
1089
 
@@ -1621,27 +1642,62 @@ function App() {
1621
  <Typography gutterBottom>CFG Scale</Typography>
1622
  <Box sx={appStyles.sliderRow}>
1623
  <Slider
1624
- value={cfgScale}
1625
  onChange={(e, value) => setCfgScale(value)}
1626
  min={0.1}
1627
  max={20}
1628
  step={0.1}
1629
  valueLabelDisplay="auto"
 
1630
  sx={appStyles.sliderFlexGrow}
1631
  />
1632
  <TextField
1633
  type="number"
1634
- value={cfgScale}
1635
  onChange={(e) => {
1636
  const val = parseFloat(e.target.value);
1637
  if (Number.isNaN(val)) return;
1638
  setCfgScale(Math.max(0.1, Math.min(20, val)));
1639
  }}
1640
  inputProps={{ min: 0.1, max: 20, step: 0.1 }}
 
1641
  sx={appStyles.sliderInputSmall}
1642
  size="small"
1643
  />
1644
  </Box>
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1645
  </Grid>
1646
 
1647
  <Grid item xs={12}>
@@ -1864,7 +1920,7 @@ function App() {
1864
  </Grid>
1865
  </TabPanel>
1866
 
1867
- <TabPanel value={tabValue} index={3}>
1868
  {performanceEnabled ? (
1869
  <Suspense fallback={
1870
  <Box sx={{ display: 'flex', justifyContent: 'center', py: 6 }}>
@@ -1872,6 +1928,7 @@ function App() {
1872
  </Box>
1873
  }>
1874
  <PerformancePanel
 
1875
  selectedModel={selectedModel}
1876
  selectedUnwrappedModel={selectedUnwrappedModel}
1877
  availableModels={availableModels}
@@ -1879,6 +1936,13 @@ function App() {
1879
  onSelectModel={setSelectedModel}
1880
  onSelectUnwrappedModel={setSelectedUnwrappedModel}
1881
  onRefreshModels={fetchAvailableModels}
 
 
 
 
 
 
 
1882
  />
1883
  </Suspense>
1884
  ) : (
 
59
  import CheckpointManager from './components/CheckpointManager';
60
  import GeneratedFragmentsWindow from './components/GeneratedFragmentsWindow';
61
  import WelcomePage from './components/WelcomePage';
62
+ import { clearPerformanceSession } from './components/usePerformanceSession';
63
  import { formatDuration } from './utils/format';
64
  import theme, { appStyles, lightTheme } from './theme';
65
 
 
138
  const [generatedFragments, setGeneratedFragments] = useState([]);
139
  const [currentFilename, setCurrentFilename] = useState('');
140
  const [cfgScale, setCfgScale] = useState(7.0);
141
+ const [steps, setSteps] = useState(250);
142
  const [batchCount, setBatchCount] = useState(1);
143
  const [randomSeed, setRandomSeed] = useState(true);
144
  const [seedValue, setSeedValue] = useState('');
 
208
  const [showStartFreshDialog, setShowStartFreshDialog] = useState(false);
209
  const [isStartingFresh, setIsStartingFresh] = useState(false);
210
  const [uploadKey, setUploadKey] = useState(0);
211
+ // Bumping this key forces the performance panel to remount, which is how
212
+ // we flush its in-memory session state on Fresh Start (clearing localStorage
213
+ // alone wouldn't reset the mounted panel's useState mirrors).
214
+ const [performanceResetKey, setPerformanceResetKey] = useState(0);
215
  const [isFreeingGPU, setIsFreeingGPU] = useState(false);
216
  const [showFreeGPUDialog, setShowFreeGPUDialog] = useState(false);
217
  const [modelWarning, setModelWarning] = useState({
 
259
  return 10;
260
  };
261
 
262
+ const isSmallModel = (() => {
263
+ if (selectedModel === 'stable-audio-open-small') return true;
264
+ const model = availableModels.find(m => m.name === selectedModel);
265
+ if (model && selectedUnwrappedModel) {
266
+ const u = model.unwrapped_models?.find(x => x.path === selectedUnwrappedModel);
267
+ return u ? (u.size_mb || 0) < 2000 : false;
268
+ }
269
+ return false;
270
+ })();
271
+
272
  useEffect(() => {
273
  const maxDuration = getMaxDuration();
274
  if (generationDuration > maxDuration) {
 
539
  const baseRequestData = {
540
  prompt: generationPrompt,
541
  duration: generationDuration,
542
+ cfg_scale: cfgScale,
543
+ steps: steps
544
  };
545
 
546
  const baseModel = baseModels.find(m => m.name === selectedModel);
 
648
  prompt: generationPrompt,
649
  duration: generationDuration,
650
  cfgScale,
651
+ steps,
652
  seed: seedForRun,
653
  batchIndex,
654
  batchTotal: totalRuns,
 
700
 
701
  const stopGeneration = () => {
702
  stopGenerationRef.current = true;
 
 
 
703
  api.post('/api/stop-generation').catch(() => {});
704
  if (generationAbortRef.current) {
705
  try { generationAbortRef.current.abort(); } catch (_) {}
 
724
  setGenerationPrompt('');
725
  setUploadKey(prev => prev + 1);
726
 
727
+ // Wipe persisted performance session and force-remount the panel so
728
+ // its in-memory state resets to defaults along with localStorage.
729
+ // (MIDI mappings and other app preferences are intentionally kept.)
730
+ clearPerformanceSession();
731
+ setPerformanceResetKey(prev => prev + 1);
732
+
733
  setProcessingStatus(response.data.message);
734
 
735
  fetchSystemStatus();
 
1101
  <Paper sx={{ p: 2 }} variant="outlined">
1102
  <Box sx={{ display: 'flex', alignItems: 'center', gap: 1, mb: 1.5 }}>
1103
  <UploadIcon size={20} />
1104
+ <Typography variant="h6">Manual Annotation</Typography>
1105
  </Box>
1106
  <Typography variant="body2" color="textSecondary" sx={{ mb: 2 }}>
1107
+ Upload audio files one by one and annotate them yourself.
1108
  Use this when you want full control over every annotation.
1109
  </Typography>
1110
 
 
1642
  <Typography gutterBottom>CFG Scale</Typography>
1643
  <Box sx={appStyles.sliderRow}>
1644
  <Slider
1645
+ value={isSmallModel ? 1.0 : cfgScale}
1646
  onChange={(e, value) => setCfgScale(value)}
1647
  min={0.1}
1648
  max={20}
1649
  step={0.1}
1650
  valueLabelDisplay="auto"
1651
+ disabled={isSmallModel}
1652
  sx={appStyles.sliderFlexGrow}
1653
  />
1654
  <TextField
1655
  type="number"
1656
+ value={isSmallModel ? 1.0 : cfgScale}
1657
  onChange={(e) => {
1658
  const val = parseFloat(e.target.value);
1659
  if (Number.isNaN(val)) return;
1660
  setCfgScale(Math.max(0.1, Math.min(20, val)));
1661
  }}
1662
  inputProps={{ min: 0.1, max: 20, step: 0.1 }}
1663
+ disabled={isSmallModel}
1664
  sx={appStyles.sliderInputSmall}
1665
  size="small"
1666
  />
1667
  </Box>
1668
+ {isSmallModel && (
1669
+ <Typography variant="caption" color="textSecondary">
1670
+ Locked at 1.0 for the distilled small model.
1671
+ </Typography>
1672
+ )}
1673
+ </Grid>
1674
+
1675
+ <Grid item xs={12}>
1676
+ <Typography gutterBottom>Inference Steps</Typography>
1677
+ <Box sx={appStyles.sliderRow}>
1678
+ <Slider
1679
+ value={steps}
1680
+ onChange={(e, value) => setSteps(value)}
1681
+ min={50}
1682
+ max={250}
1683
+ step={null}
1684
+ marks={[
1685
+ { value: 50, label: '50' },
1686
+ { value: 100, label: '100' },
1687
+ { value: 150, label: '150' },
1688
+ { value: 200, label: '200' },
1689
+ { value: 250, label: '250' },
1690
+ ]}
1691
+ valueLabelDisplay="auto"
1692
+ disabled={isSmallModel}
1693
+ sx={appStyles.sliderFlexGrow}
1694
+ />
1695
+ </Box>
1696
+ {isSmallModel && (
1697
+ <Typography variant="caption" color="textSecondary">
1698
+ Locked at 8 steps (pingpong sampler) for the distilled small model.
1699
+ </Typography>
1700
+ )}
1701
  </Grid>
1702
 
1703
  <Grid item xs={12}>
 
1920
  </Grid>
1921
  </TabPanel>
1922
 
1923
+ <TabPanel value={tabValue} index={3} keepMounted>
1924
  {performanceEnabled ? (
1925
  <Suspense fallback={
1926
  <Box sx={{ display: 'flex', justifyContent: 'center', py: 6 }}>
 
1928
  </Box>
1929
  }>
1930
  <PerformancePanel
1931
+ key={performanceResetKey}
1932
  selectedModel={selectedModel}
1933
  selectedUnwrappedModel={selectedUnwrappedModel}
1934
  availableModels={availableModels}
 
1936
  onSelectModel={setSelectedModel}
1937
  onSelectUnwrappedModel={setSelectedUnwrappedModel}
1938
  onRefreshModels={fetchAvailableModels}
1939
+ steps={steps}
1940
+ onStepsChange={setSteps}
1941
+ randomSeed={randomSeed}
1942
+ seedValue={seedValue}
1943
+ onRandomSeedChange={setRandomSeed}
1944
+ onSeedValueChange={setSeedValue}
1945
+ onPresetLoaded={() => setPerformanceResetKey(prev => prev + 1)}
1946
  />
1947
  </Suspense>
1948
  ) : (
app/frontend/src/components/BulkAnnotatePanel.js CHANGED
@@ -118,8 +118,6 @@ export default function BulkAnnotatePanel({ onCommitted }) {
118
  const resp = await api.get('/api/bulk-annotate/status');
119
  data = resp.data;
120
  } catch (exc) {
121
- // Transient errors (e.g. Flask auto-reload) must not kill polling —
122
- // the download/annotation keeps running on the backend side.
123
  return;
124
  }
125
  setStatus(data);
@@ -265,13 +263,13 @@ export default function BulkAnnotatePanel({ onCommitted }) {
265
  <Paper sx={{ p: 2, mt: 3 }} variant="outlined">
266
  <Box sx={{ display: 'flex', alignItems: 'center', gap: 1, mb: 1.5 }}>
267
  <TagsIcon size={20} />
268
- <Typography variant="h6">Bulk Auto-Annotate</Typography>
269
  </Box>
270
  <Typography variant="body2" color="textSecondary" sx={{ mb: 2 }}>
271
  Point at a folder of audio files and auto-generate prompts.
272
  Basic uses librosa (tempo + key). Rich adds CLAP tagging (genre, mood, instruments).
273
  </Typography>
274
-
275
  <Accordion sx={{ mb: 2 }}>
276
  <AccordionSummary expandIcon={<ExpandMoreIcon size={18} />}>
277
  <Box sx={{ display: 'flex', alignItems: 'center', gap: 1 }}>
 
118
  const resp = await api.get('/api/bulk-annotate/status');
119
  data = resp.data;
120
  } catch (exc) {
 
 
121
  return;
122
  }
123
  setStatus(data);
 
263
  <Paper sx={{ p: 2, mt: 3 }} variant="outlined">
264
  <Box sx={{ display: 'flex', alignItems: 'center', gap: 1, mb: 1.5 }}>
265
  <TagsIcon size={20} />
266
+ <Typography variant="h6">Bulk Auto-Annotation</Typography>
267
  </Box>
268
  <Typography variant="body2" color="textSecondary" sx={{ mb: 2 }}>
269
  Point at a folder of audio files and auto-generate prompts.
270
  Basic uses librosa (tempo + key). Rich adds CLAP tagging (genre, mood, instruments).
271
  </Typography>
272
+
273
  <Accordion sx={{ mb: 2 }}>
274
  <AccordionSummary expandIcon={<ExpandMoreIcon size={18} />}>
275
  <Box sx={{ display: 'flex', alignItems: 'center', gap: 1 }}>
app/frontend/src/components/MidiConfigMenu.js ADDED
@@ -0,0 +1,225 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React from 'react';
2
+ import {
3
+ Popover,
4
+ Box,
5
+ Typography,
6
+ FormControl,
7
+ Select,
8
+ MenuItem,
9
+ Button,
10
+ IconButton,
11
+ Tooltip,
12
+ Divider,
13
+ ToggleButton,
14
+ ToggleButtonGroup,
15
+ Alert,
16
+ } from '@mui/material';
17
+ import { Trash2 as DeleteIcon, X as CloseIcon } from 'lucide-react';
18
+ import { useMidi, formatMidi } from './MidiContext';
19
+ import { perfTokens } from '../theme';
20
+
21
+ const CHANNEL_OPTIONS = [
22
+ { value: 0, label: 'Any' },
23
+ ...Array.from({ length: 16 }, (_, i) => ({ value: i + 1, label: `Ch ${i + 1}` })),
24
+ ];
25
+
26
+ export default function MidiConfigMenu({ anchorEl, open, onClose }) {
27
+ const ctx = useMidi();
28
+ if (!ctx) return null;
29
+
30
+ const {
31
+ config,
32
+ inputs,
33
+ supported,
34
+ permissionError,
35
+ setDevice,
36
+ setChannelFilter,
37
+ setTakeover,
38
+ clearMapping,
39
+ clearAll,
40
+ } = ctx;
41
+
42
+ const sortedMappings = [...config.mappings].sort((a, b) => a.label.localeCompare(b.label));
43
+
44
+ return (
45
+ <Popover
46
+ anchorEl={anchorEl}
47
+ open={open}
48
+ onClose={onClose}
49
+ anchorOrigin={{ vertical: 'bottom', horizontal: 'right' }}
50
+ transformOrigin={{ vertical: 'top', horizontal: 'right' }}
51
+ slotProps={{
52
+ paper: {
53
+ sx: {
54
+ width: 380,
55
+ maxHeight: '70vh',
56
+ p: 2,
57
+ borderRadius: 2,
58
+ border: '1px solid',
59
+ borderColor: 'divider',
60
+ },
61
+ },
62
+ }}
63
+ >
64
+ <Box sx={{ display: 'flex', alignItems: 'center', justifyContent: 'space-between', mb: 1.5 }}>
65
+ <Typography variant="subtitle2" sx={{ letterSpacing: '0.08em', textTransform: 'uppercase', color: 'text.secondary' }}>
66
+ MIDI Settings
67
+ </Typography>
68
+ <IconButton size="small" onClick={onClose}>
69
+ <CloseIcon size={14} />
70
+ </IconButton>
71
+ </Box>
72
+
73
+ {!supported && (
74
+ <Alert severity="warning" sx={{ mb: 1.5 }}>
75
+ {permissionError || 'Web MIDI is not available in this browser. Try Chrome / Edge / Electron.'}
76
+ </Alert>
77
+ )}
78
+
79
+ <Box sx={{ display: 'flex', flexDirection: 'column', gap: 1.5 }}>
80
+ <Box>
81
+ <Typography variant="caption" sx={{ color: 'text.secondary', display: 'block', mb: 0.5 }}>
82
+ Input device
83
+ </Typography>
84
+ <FormControl size="small" fullWidth>
85
+ <Select
86
+ value={config.deviceId && inputs.some(i => i.id === config.deviceId) ? config.deviceId : ''}
87
+ onChange={(e) => setDevice(e.target.value || null)}
88
+ displayEmpty
89
+ disabled={!supported}
90
+ renderValue={(value) => {
91
+ if (!value) return <em style={{ opacity: 0.6 }}>None</em>;
92
+ const found = inputs.find(i => i.id === value);
93
+ return found ? found.name : 'Disconnected';
94
+ }}
95
+ >
96
+ <MenuItem value="">
97
+ <em>None</em>
98
+ </MenuItem>
99
+ {inputs.map((input) => (
100
+ <MenuItem key={input.id} value={input.id}>
101
+ {input.name}
102
+ </MenuItem>
103
+ ))}
104
+ </Select>
105
+ </FormControl>
106
+ {config.deviceName && !inputs.some(i => i.name === config.deviceName) && (
107
+ <Typography variant="caption" sx={{ color: 'warning.main', display: 'block', mt: 0.5 }}>
108
+ Saved device "{config.deviceName}" not connected
109
+ </Typography>
110
+ )}
111
+ </Box>
112
+
113
+ <Box sx={{ display: 'flex', gap: 1 }}>
114
+ <Box sx={{ flex: 1 }}>
115
+ <Typography variant="caption" sx={{ color: 'text.secondary', display: 'block', mb: 0.5 }}>
116
+ Channel filter
117
+ </Typography>
118
+ <FormControl size="small" fullWidth>
119
+ <Select
120
+ value={config.channelFilter}
121
+ onChange={(e) => setChannelFilter(Number(e.target.value))}
122
+ disabled={!supported}
123
+ >
124
+ {CHANNEL_OPTIONS.map(opt => (
125
+ <MenuItem key={opt.value} value={opt.value}>{opt.label}</MenuItem>
126
+ ))}
127
+ </Select>
128
+ </FormControl>
129
+ </Box>
130
+
131
+ <Box sx={{ flex: 1 }}>
132
+ <Typography variant="caption" sx={{ color: 'text.secondary', display: 'block', mb: 0.5 }}>
133
+ Takeover
134
+ </Typography>
135
+ <ToggleButtonGroup
136
+ size="small"
137
+ value={config.takeover}
138
+ exclusive
139
+ onChange={(_, v) => { if (v) setTakeover(v); }}
140
+ fullWidth
141
+ sx={{ height: 40 }}
142
+ >
143
+ <ToggleButton value="jump" sx={{ fontSize: perfTokens.fontSize.body }}>Jump</ToggleButton>
144
+ <ToggleButton value="pickup" sx={{ fontSize: perfTokens.fontSize.body }}>Pickup</ToggleButton>
145
+ </ToggleButtonGroup>
146
+ </Box>
147
+ </Box>
148
+
149
+ <Divider sx={{ my: 0.5 }} />
150
+
151
+ <Box sx={{ display: 'flex', alignItems: 'center', justifyContent: 'space-between' }}>
152
+ <Typography variant="caption" sx={{ color: 'text.secondary', letterSpacing: '0.08em', textTransform: 'uppercase' }}>
153
+ Mappings ({config.mappings.length})
154
+ </Typography>
155
+ <Button
156
+ size="small"
157
+ onClick={clearAll}
158
+ disabled={config.mappings.length === 0}
159
+ sx={{ fontSize: perfTokens.fontSize.small, textTransform: 'none' }}
160
+ >
161
+ Clear all
162
+ </Button>
163
+ </Box>
164
+
165
+ <Box
166
+ sx={{
167
+ border: '1px solid',
168
+ borderColor: 'divider',
169
+ borderRadius: 1,
170
+ maxHeight: 280,
171
+ overflowY: 'auto',
172
+ bgcolor: 'background.default',
173
+ }}
174
+ >
175
+ {sortedMappings.length === 0 ? (
176
+ <Box sx={{ p: 2, textAlign: 'center' }}>
177
+ <Typography variant="caption" sx={{ color: 'text.disabled', fontStyle: 'italic' }}>
178
+ No mappings yet. Enable MIDI mode (the MIDI button), click a control, then move a hardware knob, fader, or button.
179
+ </Typography>
180
+ </Box>
181
+ ) : (
182
+ sortedMappings.map((m) => (
183
+ <Box
184
+ key={m.controlId}
185
+ sx={{
186
+ display: 'flex',
187
+ alignItems: 'center',
188
+ gap: 1,
189
+ px: 1,
190
+ py: 0.6,
191
+ borderBottom: '1px solid',
192
+ borderColor: 'divider',
193
+ '&:last-child': { borderBottom: 'none' },
194
+ }}
195
+ >
196
+ <Box sx={{ flex: 1, minWidth: 0 }}>
197
+ <Typography variant="body2" sx={{ fontSize: perfTokens.fontSize.body, overflow: 'hidden', textOverflow: 'ellipsis', whiteSpace: 'nowrap' }}>
198
+ {m.label}
199
+ </Typography>
200
+ <Typography variant="caption" sx={{ color: 'text.secondary', fontSize: perfTokens.fontSize.small, fontFamily: 'ui-monospace, SFMono-Regular, Menlo, Consolas, monospace' }}>
201
+ {formatMidi(m.midi)}
202
+ </Typography>
203
+ </Box>
204
+ <Tooltip title="Remove mapping">
205
+ <IconButton
206
+ size="small"
207
+ onClick={() => clearMapping(m.controlId)}
208
+ sx={{ color: 'text.disabled', '&:hover': { color: 'error.main' } }}
209
+ >
210
+ <DeleteIcon size={13} />
211
+ </IconButton>
212
+ </Tooltip>
213
+ </Box>
214
+ ))
215
+ )}
216
+ </Box>
217
+
218
+ <Typography variant="caption" sx={{ color: 'text.disabled', fontSize: perfTokens.fontSize.small, lineHeight: 1.4 }}>
219
+ Pickup = ignore the hardware until its position matches the on-screen value (no jumps).
220
+ Right-click a control while in MIDI mode to clear its mapping.
221
+ </Typography>
222
+ </Box>
223
+ </Popover>
224
+ );
225
+ }
app/frontend/src/components/MidiContext.js ADDED
@@ -0,0 +1,465 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React, {
2
+ createContext,
3
+ useCallback,
4
+ useContext,
5
+ useEffect,
6
+ useMemo,
7
+ useRef,
8
+ useState,
9
+ } from 'react';
10
+ import { Box } from '@mui/material';
11
+
12
+ const STORAGE_KEY = 'fragmenta.midi.config.v1';
13
+
14
+ const DEFAULT_CONFIG = {
15
+ deviceId: null,
16
+ deviceName: null,
17
+ channelFilter: 0,
18
+ takeover: 'jump',
19
+ mappings: [],
20
+ };
21
+
22
+ const MIDI_MODE = {
23
+ NOTE_ON: 0x90,
24
+ NOTE_OFF: 0x80,
25
+ CC: 0xb0,
26
+ };
27
+
28
+ function loadConfig() {
29
+ try {
30
+ const raw = localStorage.getItem(STORAGE_KEY);
31
+ if (!raw) return { ...DEFAULT_CONFIG, mappings: [] };
32
+ const parsed = JSON.parse(raw);
33
+ return {
34
+ ...DEFAULT_CONFIG,
35
+ ...parsed,
36
+ mappings: Array.isArray(parsed.mappings) ? parsed.mappings : [],
37
+ };
38
+ } catch {
39
+ return { ...DEFAULT_CONFIG, mappings: [] };
40
+ }
41
+ }
42
+
43
+ function saveConfig(config) {
44
+ try {
45
+ localStorage.setItem(STORAGE_KEY, JSON.stringify(config));
46
+ } catch { /* quota or serialization — non-fatal */ }
47
+ }
48
+
49
+ // Wipes the persisted MIDI device + mappings. The provider will pick up the
50
+ // reset on its next mount (caller is expected to remount).
51
+ export function clearMidiConfig() {
52
+ try { localStorage.removeItem(STORAGE_KEY); }
53
+ catch { /* non-fatal */ }
54
+ }
55
+
56
+ function midiKey(midi) {
57
+ return `${midi.type}:${midi.channel}:${midi.number}`;
58
+ }
59
+
60
+ function formatMidi(midi) {
61
+ if (!midi) return '';
62
+ const t = midi.type === 'cc' ? 'CC' : 'Note';
63
+ return `${t} ${midi.number} · ch.${midi.channel}`;
64
+ }
65
+
66
+ const MidiContext = createContext(null);
67
+
68
+ export function MidiProvider({ children }) {
69
+ const [config, setConfig] = useState(loadConfig);
70
+ const [inputs, setInputs] = useState([]);
71
+ const [supported, setSupported] = useState(true);
72
+ const [permissionError, setPermissionError] = useState(null);
73
+ const [learnMode, setLearnMode] = useState(false);
74
+ const [learnTarget, setLearnTarget] = useState(null);
75
+
76
+ const accessRef = useRef(null);
77
+ const subscribersRef = useRef(new Map());
78
+ const pickupArmedRef = useRef(new Map());
79
+ const configRef = useRef(config);
80
+ const learnTargetRef = useRef(learnTarget);
81
+
82
+ useEffect(() => {
83
+ configRef.current = config;
84
+ saveConfig(config);
85
+ }, [config]);
86
+
87
+ useEffect(() => { learnTargetRef.current = learnTarget; }, [learnTarget]);
88
+
89
+ const refreshInputs = useCallback(() => {
90
+ const access = accessRef.current;
91
+ if (!access) return;
92
+ const list = [];
93
+ access.inputs.forEach((input) => {
94
+ list.push({
95
+ id: input.id,
96
+ name: input.name || 'Unknown device',
97
+ manufacturer: input.manufacturer || '',
98
+ });
99
+ });
100
+ setInputs(list);
101
+ }, []);
102
+
103
+ useEffect(() => {
104
+ if (typeof navigator === 'undefined' || !navigator.requestMIDIAccess) {
105
+ setSupported(false);
106
+ return undefined;
107
+ }
108
+ let cancelled = false;
109
+ navigator.requestMIDIAccess({ sysex: false })
110
+ .then((access) => {
111
+ if (cancelled) return;
112
+ accessRef.current = access;
113
+ refreshInputs();
114
+ access.onstatechange = refreshInputs;
115
+ })
116
+ .catch((err) => {
117
+ setPermissionError(err?.message || 'MIDI permission denied');
118
+ setSupported(false);
119
+ });
120
+ return () => { cancelled = true; };
121
+ }, [refreshInputs]);
122
+
123
+ useEffect(() => {
124
+ if (!inputs.length || !config.deviceName) return;
125
+ const stillThere = config.deviceId && inputs.some(i => i.id === config.deviceId);
126
+ if (stillThere) return;
127
+ const byName = inputs.find(i => i.name === config.deviceName);
128
+ if (byName) {
129
+ setConfig(prev => ({ ...prev, deviceId: byName.id }));
130
+ }
131
+ }, [inputs, config.deviceId, config.deviceName]);
132
+
133
+ const captureLearn = useCallback((controlId, midi) => {
134
+ setConfig((prev) => {
135
+ const subOpts = subscribersRef.current.get(controlId)?.opts || {};
136
+ const newMapping = {
137
+ controlId,
138
+ label: subOpts.label || controlId,
139
+ kind: subOpts.kind || 'continuous',
140
+ curve: subOpts.curve || 'linear',
141
+ min: subOpts.min ?? 0,
142
+ max: subOpts.max ?? 1,
143
+ midi,
144
+ };
145
+ const targetKey = midiKey(midi);
146
+ const filtered = prev.mappings.filter(
147
+ (m) => m.controlId !== controlId && midiKey(m.midi) !== targetKey,
148
+ );
149
+ return { ...prev, mappings: [...filtered, newMapping] };
150
+ });
151
+ pickupArmedRef.current.delete(controlId);
152
+ setLearnTarget(null);
153
+ }, []);
154
+
155
+ const dispatchMessage = useCallback((event) => {
156
+ const data = event.data;
157
+ if (!data || data.length < 2) return;
158
+ const status = data[0];
159
+ const data1 = data[1];
160
+ const data2 = data.length > 2 ? data[2] : 0;
161
+ const type = status & 0xf0;
162
+ const channel = (status & 0x0f) + 1;
163
+ const cfg = configRef.current;
164
+
165
+ if (cfg.channelFilter && channel !== cfg.channelFilter) return;
166
+
167
+ const isCC = type === MIDI_MODE.CC;
168
+ const isNoteOn = type === MIDI_MODE.NOTE_ON && data2 > 0;
169
+ const isNoteOff = type === MIDI_MODE.NOTE_OFF || (type === MIDI_MODE.NOTE_ON && data2 === 0);
170
+ if (!isCC && !isNoteOn && !isNoteOff) return;
171
+
172
+ const incomingType = isCC ? 'cc' : 'note';
173
+ const target = learnTargetRef.current;
174
+ if (target && (isCC || isNoteOn)) {
175
+ captureLearn(target, { type: incomingType, channel, number: data1 });
176
+ return;
177
+ }
178
+
179
+ for (const m of cfg.mappings) {
180
+ if (m.midi.channel !== channel || m.midi.number !== data1) continue;
181
+ if (m.midi.type !== incomingType) continue;
182
+ const sub = subscribersRef.current.get(m.controlId);
183
+ if (!sub) continue;
184
+
185
+ if (sub.opts.kind === 'continuous') {
186
+ if (!isCC) continue;
187
+ applyContinuous(sub, m, data2, cfg.takeover);
188
+ } else if (sub.opts.kind === 'trigger') {
189
+ if (isNoteOn) sub.handler();
190
+ else if (isCC && data2 >= 64) sub.handler();
191
+ }
192
+ }
193
+ }, [captureLearn]);
194
+
195
+ useEffect(() => {
196
+ const access = accessRef.current;
197
+ if (!access) return undefined;
198
+ const bound = [];
199
+ access.inputs.forEach((input) => {
200
+ if (config.deviceId && input.id === config.deviceId) {
201
+ input.onmidimessage = dispatchMessage;
202
+ bound.push(input);
203
+ } else {
204
+ input.onmidimessage = null;
205
+ }
206
+ });
207
+
208
+ pickupArmedRef.current = new Map();
209
+ return () => {
210
+ bound.forEach((i) => { i.onmidimessage = null; });
211
+ };
212
+ }, [config.deviceId, inputs, dispatchMessage]);
213
+
214
+ function applyContinuous(sub, mapping, midiValue, takeover) {
215
+ const norm = midiValue / 127;
216
+ let target;
217
+ if (mapping.curve === 'log' && mapping.min > 0 && mapping.max > 0) {
218
+ target = mapping.min * Math.pow(mapping.max / mapping.min, norm);
219
+ } else {
220
+ target = mapping.min + norm * (mapping.max - mapping.min);
221
+ }
222
+
223
+ if (takeover === 'pickup') {
224
+ const armed = pickupArmedRef.current.get(mapping.controlId);
225
+ if (!armed) {
226
+ const current = typeof sub.getValue === 'function' ? sub.getValue() : sub.value;
227
+ const span = mapping.max - mapping.min;
228
+ if (span === 0 || !isFinite(current)) {
229
+ pickupArmedRef.current.set(mapping.controlId, true);
230
+ } else {
231
+ // Compare on the same curve we used to compute target.
232
+ let currentNorm;
233
+ if (mapping.curve === 'log' && mapping.min > 0 && current > 0) {
234
+ currentNorm = Math.log(current / mapping.min) / Math.log(mapping.max / mapping.min);
235
+ } else {
236
+ currentNorm = (current - mapping.min) / span;
237
+ }
238
+ if (Math.abs(norm - currentNorm) < 0.02) {
239
+ pickupArmedRef.current.set(mapping.controlId, true);
240
+ } else {
241
+ return;
242
+ }
243
+ }
244
+ }
245
+ }
246
+ sub.handler(target);
247
+ }
248
+
249
+ const beginLearn = useCallback((controlId) => {
250
+ setLearnMode(true);
251
+ setLearnTarget(controlId);
252
+ }, []);
253
+
254
+ const cancelLearn = useCallback(() => setLearnTarget(null), []);
255
+
256
+ const clearMapping = useCallback((controlId) => {
257
+ setConfig((prev) => ({
258
+ ...prev,
259
+ mappings: prev.mappings.filter((m) => m.controlId !== controlId),
260
+ }));
261
+ pickupArmedRef.current.delete(controlId);
262
+ }, []);
263
+
264
+ const clearAll = useCallback(() => {
265
+ setConfig((prev) => ({ ...prev, mappings: [] }));
266
+ pickupArmedRef.current.clear();
267
+ }, []);
268
+
269
+ const setDevice = useCallback((deviceId) => {
270
+ const found = inputs.find((i) => i.id === deviceId);
271
+ setConfig((prev) => ({
272
+ ...prev,
273
+ deviceId: deviceId || null,
274
+ deviceName: found ? found.name : null,
275
+ }));
276
+ }, [inputs]);
277
+
278
+ const setChannelFilter = useCallback((channel) => {
279
+ setConfig((prev) => ({ ...prev, channelFilter: channel }));
280
+ }, []);
281
+
282
+ const setTakeover = useCallback((mode) => {
283
+ setConfig((prev) => ({ ...prev, takeover: mode }));
284
+ pickupArmedRef.current.clear();
285
+ }, []);
286
+
287
+ const exitLearnMode = useCallback(() => {
288
+ setLearnMode(false);
289
+ setLearnTarget(null);
290
+ }, []);
291
+
292
+ const toggleLearnMode = useCallback(() => {
293
+ setLearnMode((prev) => {
294
+ if (prev) setLearnTarget(null);
295
+ return !prev;
296
+ });
297
+ }, []);
298
+
299
+ const registerSubscriber = useCallback((id, sub) => {
300
+ subscribersRef.current.set(id, sub);
301
+ }, []);
302
+
303
+ const unregisterSubscriber = useCallback((id) => {
304
+ subscribersRef.current.delete(id);
305
+ }, []);
306
+
307
+ const value = useMemo(() => ({
308
+ config,
309
+ inputs,
310
+ supported,
311
+ permissionError,
312
+ learnMode,
313
+ learnTarget,
314
+ setDevice,
315
+ setChannelFilter,
316
+ setTakeover,
317
+ beginLearn,
318
+ cancelLearn,
319
+ clearMapping,
320
+ clearAll,
321
+ toggleLearnMode,
322
+ exitLearnMode,
323
+ registerSubscriber,
324
+ unregisterSubscriber,
325
+ }), [
326
+ config, inputs, supported, permissionError, learnMode, learnTarget,
327
+ setDevice, setChannelFilter, setTakeover, beginLearn, cancelLearn,
328
+ clearMapping, clearAll, toggleLearnMode, exitLearnMode,
329
+ registerSubscriber, unregisterSubscriber,
330
+ ]);
331
+
332
+ return <MidiContext.Provider value={value}>{children}</MidiContext.Provider>;
333
+ }
334
+
335
+ export function useMidi() {
336
+ const ctx = useContext(MidiContext);
337
+ return ctx;
338
+ }
339
+
340
+ export { formatMidi };
341
+
342
+
343
+ export function MidiMappable({
344
+ id,
345
+ label,
346
+ kind = 'continuous',
347
+ curve = 'linear',
348
+ min = 0,
349
+ max = 1,
350
+ value,
351
+ onChange,
352
+ sx,
353
+ children,
354
+ }) {
355
+ const ctx = useMidi();
356
+ const valueRef = useRef(value);
357
+ useEffect(() => { valueRef.current = value; }, [value]);
358
+ const handlerRef = useRef(onChange);
359
+ useEffect(() => { handlerRef.current = onChange; }, [onChange]);
360
+
361
+ useEffect(() => {
362
+ if (!ctx) return undefined;
363
+ ctx.registerSubscriber(id, {
364
+ opts: { kind, curve, min, max, label },
365
+ handler: (v) => handlerRef.current?.(v),
366
+ getValue: () => valueRef.current,
367
+ });
368
+ return () => ctx.unregisterSubscriber(id);
369
+ }, [ctx, id, kind, curve, min, max, label]);
370
+
371
+ if (!ctx) return <>{children}</>;
372
+
373
+ const mapping = ctx.config.mappings.find((m) => m.controlId === id);
374
+ const isLearningThis = ctx.learnTarget === id;
375
+ const showOverlay = ctx.learnMode;
376
+
377
+ return (
378
+ <Box sx={{ position: 'relative', display: 'flex', flexDirection: 'column', minWidth: 0, minHeight: 0, ...sx }}>
379
+ {children}
380
+ {mapping && !showOverlay && (
381
+ <Box
382
+ sx={{
383
+ position: 'absolute',
384
+ top: 1,
385
+ right: 1,
386
+ fontSize: '0.5rem',
387
+ fontFamily: 'ui-monospace, SFMono-Regular, Menlo, Consolas, monospace',
388
+ bgcolor: 'rgba(83, 193, 138, 0.85)',
389
+ color: '#000',
390
+ px: 0.4,
391
+ py: 0.05,
392
+ borderRadius: 0.5,
393
+ pointerEvents: 'none',
394
+ letterSpacing: '0.04em',
395
+ zIndex: 5,
396
+ }}
397
+ >
398
+ {mapping.midi.type === 'cc' ? 'CC' : 'N'}{mapping.midi.number}
399
+ </Box>
400
+ )}
401
+ {showOverlay && (
402
+ <Box
403
+ onClick={(e) => {
404
+ e.preventDefault();
405
+ e.stopPropagation();
406
+ if (isLearningThis) ctx.cancelLearn();
407
+ else ctx.beginLearn(id);
408
+ }}
409
+ onContextMenu={(e) => {
410
+ e.preventDefault();
411
+ e.stopPropagation();
412
+ if (mapping) ctx.clearMapping(id);
413
+ }}
414
+ title={
415
+ isLearningThis
416
+ ? `${label}: move a hardware control to bind (right-click to clear)`
417
+ : mapping
418
+ ? `${label}: ${formatMidi(mapping.midi)} (click to re-learn, right-click to clear)`
419
+ : `${label}: click then move a hardware control to bind`
420
+ }
421
+ sx={{
422
+ position: 'absolute',
423
+ inset: 0,
424
+ cursor: 'pointer',
425
+ zIndex: 20,
426
+ bgcolor: isLearningThis
427
+ ? 'rgba(245, 197, 66, 0.32)'
428
+ : mapping
429
+ ? 'rgba(83, 193, 138, 0.18)'
430
+ : 'rgba(245, 197, 66, 0.10)',
431
+ border: '1px dashed',
432
+ borderColor: isLearningThis
433
+ ? '#F5C542'
434
+ : mapping
435
+ ? 'rgba(83, 193, 138, 0.7)'
436
+ : 'rgba(245, 197, 66, 0.65)',
437
+ borderRadius: 1,
438
+ animation: isLearningThis ? 'midiPulse 900ms ease-in-out infinite' : 'none',
439
+ '@keyframes midiPulse': {
440
+ '0%, 100%': { opacity: 0.5 },
441
+ '50%': { opacity: 1 },
442
+ },
443
+ }}
444
+ >
445
+ {(mapping || isLearningThis) && (
446
+ <Box
447
+ sx={{
448
+ position: 'absolute',
449
+ top: 2,
450
+ left: 2,
451
+ fontSize: '0.5rem',
452
+ fontFamily: 'ui-monospace, SFMono-Regular, Menlo, Consolas, monospace',
453
+ color: isLearningThis ? '#F5C542' : 'rgba(83, 193, 138, 0.95)',
454
+ letterSpacing: '0.04em',
455
+ pointerEvents: 'none',
456
+ }}
457
+ >
458
+ {isLearningThis ? 'learn…' : `${mapping.midi.type === 'cc' ? 'CC' : 'N'}${mapping.midi.number}`}
459
+ </Box>
460
+ )}
461
+ </Box>
462
+ )}
463
+ </Box>
464
+ );
465
+ }
app/frontend/src/components/PerformanceChannel.js CHANGED
@@ -19,24 +19,32 @@ import {
19
  Volume2 as VolumeIcon,
20
  VolumeX as MuteIcon,
21
  } from 'lucide-react';
22
- import { performanceChannelStyles as styles } from '../theme';
 
23
 
24
  const CHANNEL_COLORS = [
25
  '#35C2D4', '#9F8AE6', '#53C18A', '#E3A34B',
26
  '#E36C61', '#F08AD2', '#5BA0F0', '#A8D86B',
27
  ];
28
 
29
- // -6 dBFS in linear amplitude — must match DEFAULT_CHANNEL_GAIN in performanceAudio.js
30
- const DEFAULT_GAIN = Math.pow(10, -6 / 20);
 
 
 
 
 
31
 
32
  const KNOB_DEFS = [
33
- { key: 'gain', label: 'GAIN', min: 0, max: 1.5, step: 0.01, default: DEFAULT_GAIN },
34
- { key: 'filter', label: 'LPF', min: 200, max: 18000, step: 1, default: 18000, log: true },
 
 
 
35
  { key: 'delay', label: 'DLY', min: 0, max: 1.0, step: 0.01, default: 0.0 },
36
  { key: 'reverb', label: 'REV', min: 0, max: 1.0, step: 0.01, default: 0.0 },
37
  ];
38
 
39
- // Snap pan to center when within this fraction of full travel from 0.
40
  const PAN_CENTER_SNAP = 0.06;
41
 
42
  const BARS_OPTIONS = [1, 2, 4, 8, 16];
@@ -45,10 +53,14 @@ const BEATS_PER_BAR = 4;
45
  export default function PerformanceChannel({
46
  index,
47
  strip,
 
 
48
  onGenerate,
49
  canGenerate,
50
  onMuteSoloChange,
51
  onStateChange,
 
 
52
  maxDuration = 47,
53
  bpm = 120,
54
  }) {
@@ -57,28 +69,43 @@ export default function PerformanceChannel({
57
  const meterRef = useRef(null);
58
  const meterRafRef = useRef(null);
59
 
60
- const [prompt, setPrompt] = useState('');
61
- const [duration, setDuration] = useState(8);
62
- const [durationMode, setDurationMode] = useState('seconds');
63
- const [bars, setBars] = useState(4);
 
 
 
 
 
 
 
 
64
  const [generating, setGenerating] = useState(false);
65
  const [loaded, setLoaded] = useState(false);
66
- const [playing, setPlaying] = useState(false);
67
- const [looping, setLooping] = useState(true);
68
- const [muted, setMuted] = useState(false);
69
- const [soloed, setSoloed] = useState(false);
70
- const [knobs, setKnobs] = useState(() => {
71
- const initial = Object.fromEntries(KNOB_DEFS.map(k => [k.key, k.default]));
72
- initial.pan = 0;
73
- return initial;
74
- });
 
 
 
 
 
 
 
 
75
 
76
  const secondsFromBars = useMemo(
77
  () => bars * (60 / Math.max(bpm, 1)) * BEATS_PER_BAR,
78
  [bars, bpm]
79
  );
80
 
81
- // Only show bar counts whose implied duration fits within the model's max.
82
  const availableBars = useMemo(() => {
83
  const maxBars = (maxDuration * bpm) / (60 * BEATS_PER_BAR);
84
  const opts = BARS_OPTIONS.filter(b => b <= maxBars);
@@ -108,6 +135,28 @@ export default function PerformanceChannel({
108
 
109
  useEffect(() => { drawWave(); }, [drawWave, loaded]);
110
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
111
  useEffect(() => {
112
  setDuration(prev => Math.min(prev, maxDuration));
113
  }, [maxDuration]);
@@ -120,10 +169,17 @@ export default function PerformanceChannel({
120
 
121
  const handleGenerate = async () => {
122
  if (!prompt.trim() || generating) return;
123
- const effectiveDuration = durationMode === 'bars' ? secondsFromBars : duration;
 
124
  setGenerating(true);
125
  try {
126
- const blob = await onGenerate({ prompt, duration: effectiveDuration });
 
 
 
 
 
 
127
  await strip.loadBlob(blob);
128
  setLoaded(true);
129
  onStateChange?.(index, { loaded: true });
@@ -137,14 +193,13 @@ export default function PerformanceChannel({
137
 
138
  const handlePlay = () => {
139
  if (!loaded) return;
140
- strip.play(looping);
141
- setPlaying(true);
142
  onStateChange?.(index, { playing: true });
143
  };
144
 
145
  const handleStop = () => {
146
  strip.stop();
147
- setPlaying(false);
148
  onStateChange?.(index, { playing: false });
149
  };
150
 
@@ -170,24 +225,42 @@ export default function PerformanceChannel({
170
 
171
  const handleKnob = (key, value) => {
172
  setKnobs(prev => ({ ...prev, [key]: value }));
173
- if (key === 'gain') strip.setUserGain(value);
174
  else if (key === 'pan') strip.setPan(value);
175
  else if (key === 'filter') strip.setFilter(value);
176
  else if (key === 'delay') strip.setDelayMix(value);
177
  else if (key === 'reverb') strip.setReverbMix(value);
178
  };
179
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
180
  return (
181
  <Box sx={styles.strip(color, playing)}>
182
  <Box sx={styles.stripHeader(color)}>
183
  <Box sx={styles.channelBadge(color)}>{String(index + 1).padStart(2, '0')}</Box>
184
  <Box sx={styles.muteSoloRow}>
185
- <Tooltip title="Mute">
186
- <IconButton size="small" onClick={handleMuteToggle} sx={styles.muteBtn(muted)}>M</IconButton>
187
- </Tooltip>
188
- <Tooltip title="Solo">
189
- <IconButton size="small" onClick={handleSoloToggle} sx={styles.soloBtn(soloed)}>S</IconButton>
190
- </Tooltip>
 
 
 
 
191
  </Box>
192
  </Box>
193
 
@@ -204,9 +277,6 @@ export default function PerformanceChannel({
204
  sx={styles.promptField}
205
  disabled={generating}
206
  />
207
- {/* Fixed height keeps the strip stable when swapping between
208
- sec (slider) and bars (Select) modes — the Select would
209
- otherwise be taller and push everything below it up/down. */}
210
  <Box sx={{ ...styles.durationRow, minHeight: 26, height: 26 }}>
211
  <Box
212
  sx={{
@@ -226,8 +296,8 @@ export default function PerformanceChannel({
226
  key={mode}
227
  onClick={() => setDurationMode(value)}
228
  sx={{
229
- fontSize: '0.58rem',
230
- letterSpacing: '0.08em',
231
  textTransform: 'uppercase',
232
  fontFamily: 'inherit',
233
  px: 0.7,
@@ -268,7 +338,7 @@ export default function PerformanceChannel({
268
  size="small"
269
  sx={{
270
  flex: 1,
271
- fontSize: '0.7rem',
272
  height: '100%',
273
  '& .MuiOutlinedInput-input': {
274
  py: 0,
@@ -283,21 +353,23 @@ export default function PerformanceChannel({
283
  }}
284
  >
285
  {availableBars.map(b => (
286
- <MenuItem key={b} value={b} sx={{ fontSize: '0.75rem' }}>
287
  {b} {b === 1 ? 'bar' : 'bars'}
288
  </MenuItem>
289
  ))}
290
  </Select>
291
  )}
292
  </Box>
293
- <IconButton
294
- onClick={handleGenerate}
295
- disabled={!canGenerate || !prompt.trim() || generating}
296
- sx={styles.generateBtn(color)}
297
- size="small"
298
- >
299
- {generating ? <CircularProgress size={16} sx={{ color }} /> : <GenerateIcon size={16} />}
300
- </IconButton>
 
 
301
  </Box>
302
 
303
  <Box sx={styles.waveformWrap}>
@@ -316,72 +388,108 @@ export default function PerformanceChannel({
316
 
317
  <Box sx={{ px: 1, py: 1 }}>
318
  <Box sx={{ display: 'flex', alignItems: 'center', gap: 1, mb: 1 }}>
319
- <Box component="span" sx={{ fontSize: '0.53rem', color: 'text.secondary', letterSpacing: '0.06em', minWidth: 28 }}>PAN</Box>
320
- <Slider
321
- value={knobs.pan ?? 0}
322
- onChange={(_, v) => {
323
- // Snap to center when close, so "0" isn't fiddly to hit.
324
- const snapped = Math.abs(v) < PAN_CENTER_SNAP ? 0 : v;
325
- handleKnob('pan', snapped);
326
- }}
327
  min={-1}
328
  max={1}
329
- step={0.01}
330
- size="small"
331
- track={false}
332
- marks={[{ value: 0 }]}
333
- sx={{
334
- flex: 1,
335
- '& .MuiSlider-mark': {
336
- width: 2,
337
- height: 10,
338
- borderRadius: 1,
339
- backgroundColor: 'text.secondary',
340
- opacity: 0.8,
341
- },
342
- '& .MuiSlider-markActive': {
343
- backgroundColor: 'text.secondary',
344
- opacity: 0.8,
345
- },
346
- }}
347
- />
 
 
 
 
 
 
 
 
 
 
348
  </Box>
349
  </Box>
350
 
351
  <Box sx={styles.knobsGrid}>
352
- {KNOB_DEFS.map((k) => (
353
- <Box key={k.key} sx={styles.knobCell}>
354
- <Slider
355
- orientation="vertical"
356
- value={knobs[k.key]}
357
- onChange={(_, v) => handleKnob(k.key, v)}
358
- min={k.min}
359
- max={k.max}
360
- step={k.step}
361
- size="small"
362
- sx={styles.knobSlider(color)}
363
- />
364
- <Box component="span" sx={styles.knobLabel}>{k.label}</Box>
365
- </Box>
366
- ))}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
367
  </Box>
368
 
369
  <Box sx={styles.transportRow}>
370
- <IconButton
371
- onClick={playing ? handleStop : handlePlay}
372
- disabled={!loaded}
373
- sx={styles.transportBtn(color, playing)}
374
- size="small"
375
- >
376
- {playing ? <StopIcon size={16} /> : <PlayIcon size={16} />}
377
- </IconButton>
378
- <IconButton
379
- onClick={handleLoopToggle}
380
- sx={styles.loopBtn(color, looping)}
381
- size="small"
382
- >
383
- <LoopIcon size={14} />
384
- </IconButton>
 
 
 
 
385
  <Box sx={styles.meterTrack}>
386
  <Box ref={meterRef} sx={styles.meterFill(color)} />
387
  </Box>
 
19
  Volume2 as VolumeIcon,
20
  VolumeX as MuteIcon,
21
  } from 'lucide-react';
22
+ import { performanceChannelStyles as styles, perfTokens } from '../theme';
23
+ import { MidiMappable } from './MidiContext';
24
 
25
  const CHANNEL_COLORS = [
26
  '#35C2D4', '#9F8AE6', '#53C18A', '#E3A34B',
27
  '#E36C61', '#F08AD2', '#5BA0F0', '#A8D86B',
28
  ];
29
 
30
+ // Channel gain runs on the same dBFS scale as the master fader so the two
31
+ // scales line up: -60 dB floor, 0 dB ceiling, default at -6 dB. The knob's
32
+ // dB value is converted to linear before reaching the audio graph.
33
+ const GAIN_DB_MIN = -60;
34
+ const GAIN_DB_MAX = 0;
35
+ const GAIN_DB_DEFAULT = -6;
36
+ const gainDbToLinear = (db) => (db <= GAIN_DB_MIN ? 0 : Math.pow(10, db / 20));
37
 
38
  const KNOB_DEFS = [
39
+ { key: 'gain', label: 'GAIN', min: GAIN_DB_MIN, max: GAIN_DB_MAX, step: 0.5, default: GAIN_DB_DEFAULT },
40
+ // LPF range goes from 20 Hz (full kill) to 20 kHz (bypass). We render the
41
+ // slider on a log axis so each octave gets equal travel — without this
42
+ // the bottom 5% of the knob does all the audible work.
43
+ { key: 'filter', label: 'LPF', min: 20, max: 20000, step: 1, default: 20000, scale: 'log' },
44
  { key: 'delay', label: 'DLY', min: 0, max: 1.0, step: 0.01, default: 0.0 },
45
  { key: 'reverb', label: 'REV', min: 0, max: 1.0, step: 0.01, default: 0.0 },
46
  ];
47
 
 
48
  const PAN_CENTER_SNAP = 0.06;
49
 
50
  const BARS_OPTIONS = [1, 2, 4, 8, 16];
 
53
  export default function PerformanceChannel({
54
  index,
55
  strip,
56
+ engine,
57
+ playing = false,
58
  onGenerate,
59
  canGenerate,
60
  onMuteSoloChange,
61
  onStateChange,
62
+ onFormStateChange,
63
+ initialFormState,
64
  maxDuration = 47,
65
  bpm = 120,
66
  }) {
 
69
  const meterRef = useRef(null);
70
  const meterRafRef = useRef(null);
71
 
72
+ const init = initialFormState || {};
73
+ const initKnobs = init.knobs || {};
74
+ const defaultKnobs = (() => {
75
+ const d = Object.fromEntries(KNOB_DEFS.map(k => [k.key, k.default]));
76
+ d.pan = 0;
77
+ return d;
78
+ })();
79
+
80
+ const [prompt, setPrompt] = useState(init.prompt ?? '');
81
+ const [duration, setDuration] = useState(init.duration ?? 8);
82
+ const [durationMode, setDurationMode] = useState(init.durationMode ?? 'seconds');
83
+ const [bars, setBars] = useState(init.bars ?? 4);
84
  const [generating, setGenerating] = useState(false);
85
  const [loaded, setLoaded] = useState(false);
86
+ const [looping, setLooping] = useState(init.looping ?? true);
87
+ const [muted, setMuted] = useState(init.muted ?? false);
88
+ const [soloed, setSoloed] = useState(init.soloed ?? false);
89
+ const [knobs, setKnobs] = useState(() => ({ ...defaultKnobs, ...initKnobs }));
90
+
91
+ // Mirror form state up to the panel so it can persist the session. Skip the
92
+ // first render so we don't re-write what we just loaded from localStorage.
93
+ const initialReportSkippedRef = useRef(false);
94
+ useEffect(() => {
95
+ if (!initialReportSkippedRef.current) {
96
+ initialReportSkippedRef.current = true;
97
+ return;
98
+ }
99
+ onFormStateChange?.(index, {
100
+ prompt, duration, durationMode, bars, looping, muted, soloed, knobs,
101
+ });
102
+ }, [prompt, duration, durationMode, bars, looping, muted, soloed, knobs, index, onFormStateChange]);
103
 
104
  const secondsFromBars = useMemo(
105
  () => bars * (60 / Math.max(bpm, 1)) * BEATS_PER_BAR,
106
  [bars, bpm]
107
  );
108
 
 
109
  const availableBars = useMemo(() => {
110
  const maxBars = (maxDuration * bpm) / (60 * BEATS_PER_BAR);
111
  const opts = BARS_OPTIONS.filter(b => b <= maxBars);
 
135
 
136
  useEffect(() => { drawWave(); }, [drawWave, loaded]);
137
 
138
+ // One-shot: push restored knob/loop values into the audio strip when it
139
+ // first becomes available, so the persisted session matches what's heard.
140
+ // Mute/solo applies through the parent's mix handler so the panel can
141
+ // recompute the "any-soloed" cross-channel state.
142
+ const stripStateAppliedRef = useRef(false);
143
+ useEffect(() => {
144
+ if (!strip || stripStateAppliedRef.current) return;
145
+ stripStateAppliedRef.current = true;
146
+ strip.setUserGain(gainDbToLinear(knobs.gain));
147
+ strip.setPan(knobs.pan);
148
+ strip.setFilter(knobs.filter);
149
+ strip.setDelayMix(knobs.delay);
150
+ strip.setReverbMix(knobs.reverb);
151
+ strip.setLoop(looping);
152
+ if (muted || soloed) {
153
+ onMuteSoloChange?.(index, { mute: muted, solo: soloed });
154
+ }
155
+ // Initial values are intentionally only applied once; subsequent edits
156
+ // flow through the normal handlers below.
157
+ // eslint-disable-next-line react-hooks/exhaustive-deps
158
+ }, [strip]);
159
+
160
  useEffect(() => {
161
  setDuration(prev => Math.min(prev, maxDuration));
162
  }, [maxDuration]);
 
169
 
170
  const handleGenerate = async () => {
171
  if (!prompt.trim() || generating) return;
172
+ const inBarsMode = durationMode === 'bars';
173
+ const effectiveDuration = inBarsMode ? secondsFromBars : duration;
174
  setGenerating(true);
175
  try {
176
+ const blob = await onGenerate({
177
+ prompt,
178
+ duration: effectiveDuration,
179
+ // Only forward alignment params in bars mode — seconds mode
180
+ // generates raw audio with no post-processing.
181
+ ...(inBarsMode ? { alignBars: bars, alignBpm: bpm } : {}),
182
+ });
183
  await strip.loadBlob(blob);
184
  setLoaded(true);
185
  onStateChange?.(index, { loaded: true });
 
193
 
194
  const handlePlay = () => {
195
  if (!loaded) return;
196
+ if (engine) engine.playChannel(index, looping);
197
+ else strip.play(looping);
198
  onStateChange?.(index, { playing: true });
199
  };
200
 
201
  const handleStop = () => {
202
  strip.stop();
 
203
  onStateChange?.(index, { playing: false });
204
  };
205
 
 
225
 
226
  const handleKnob = (key, value) => {
227
  setKnobs(prev => ({ ...prev, [key]: value }));
228
+ if (key === 'gain') strip.setUserGain(gainDbToLinear(value));
229
  else if (key === 'pan') strip.setPan(value);
230
  else if (key === 'filter') strip.setFilter(value);
231
  else if (key === 'delay') strip.setDelayMix(value);
232
  else if (key === 'reverb') strip.setReverbMix(value);
233
  };
234
 
235
+ const handlePan = (v) => {
236
+ const snapped = Math.abs(v) < PAN_CENTER_SNAP ? 0 : v;
237
+ handleKnob('pan', snapped);
238
+ };
239
+
240
+ const handleTransportToggle = () => {
241
+ if (!loaded) return;
242
+ if (playing) handleStop();
243
+ else handlePlay();
244
+ };
245
+
246
+ const ctrlId = (suffix) => `channel.${index}.${suffix}`;
247
+ const ctrlLabel = (name) => `Ch ${index + 1} · ${name}`;
248
+
249
  return (
250
  <Box sx={styles.strip(color, playing)}>
251
  <Box sx={styles.stripHeader(color)}>
252
  <Box sx={styles.channelBadge(color)}>{String(index + 1).padStart(2, '0')}</Box>
253
  <Box sx={styles.muteSoloRow}>
254
+ <MidiMappable id={ctrlId('mute')} label={ctrlLabel('Mute')} kind="trigger" onChange={handleMuteToggle}>
255
+ <Tooltip title="Mute">
256
+ <IconButton size="small" onClick={handleMuteToggle} sx={styles.muteBtn(muted)}>M</IconButton>
257
+ </Tooltip>
258
+ </MidiMappable>
259
+ <MidiMappable id={ctrlId('solo')} label={ctrlLabel('Solo')} kind="trigger" onChange={handleSoloToggle}>
260
+ <Tooltip title="Solo">
261
+ <IconButton size="small" onClick={handleSoloToggle} sx={styles.soloBtn(soloed)}>S</IconButton>
262
+ </Tooltip>
263
+ </MidiMappable>
264
  </Box>
265
  </Box>
266
 
 
277
  sx={styles.promptField}
278
  disabled={generating}
279
  />
 
 
 
280
  <Box sx={{ ...styles.durationRow, minHeight: 26, height: 26 }}>
281
  <Box
282
  sx={{
 
296
  key={mode}
297
  onClick={() => setDurationMode(value)}
298
  sx={{
299
+ fontSize: perfTokens.fontSize.small,
300
+ letterSpacing: perfTokens.letterSpacing.wide,
301
  textTransform: 'uppercase',
302
  fontFamily: 'inherit',
303
  px: 0.7,
 
338
  size="small"
339
  sx={{
340
  flex: 1,
341
+ fontSize: perfTokens.fontSize.body,
342
  height: '100%',
343
  '& .MuiOutlinedInput-input': {
344
  py: 0,
 
353
  }}
354
  >
355
  {availableBars.map(b => (
356
+ <MenuItem key={b} value={b} sx={{ fontSize: perfTokens.fontSize.body }}>
357
  {b} {b === 1 ? 'bar' : 'bars'}
358
  </MenuItem>
359
  ))}
360
  </Select>
361
  )}
362
  </Box>
363
+ <MidiMappable id={ctrlId('generate')} label={ctrlLabel('Generate')} kind="trigger" onChange={handleGenerate}>
364
+ <IconButton
365
+ onClick={handleGenerate}
366
+ disabled={!canGenerate || !prompt.trim() || generating}
367
+ sx={styles.generateBtn(color)}
368
+ size="small"
369
+ >
370
+ {generating ? <CircularProgress size={16} sx={{ color }} /> : <GenerateIcon size={16} />}
371
+ </IconButton>
372
+ </MidiMappable>
373
  </Box>
374
 
375
  <Box sx={styles.waveformWrap}>
 
388
 
389
  <Box sx={{ px: 1, py: 1 }}>
390
  <Box sx={{ display: 'flex', alignItems: 'center', gap: 1, mb: 1 }}>
391
+ <Box component="span" sx={{ fontSize: perfTokens.fontSize.knob, color: 'text.secondary', letterSpacing: perfTokens.letterSpacing.wide, minWidth: 28 }}>PAN</Box>
392
+ <MidiMappable
393
+ id={ctrlId('pan')}
394
+ label={ctrlLabel('Pan')}
395
+ kind="continuous"
 
 
 
396
  min={-1}
397
  max={1}
398
+ value={knobs.pan ?? 0}
399
+ onChange={handlePan}
400
+ sx={{ flex: 1, flexDirection: 'row' }}
401
+ >
402
+ <Slider
403
+ value={knobs.pan ?? 0}
404
+ onChange={(_, v) => handlePan(v)}
405
+ min={-1}
406
+ max={1}
407
+ step={0.01}
408
+ size="small"
409
+ track={false}
410
+ marks={[{ value: 0 }]}
411
+ sx={{
412
+ flex: 1,
413
+ '& .MuiSlider-mark': {
414
+ width: 2,
415
+ height: 10,
416
+ borderRadius: 1,
417
+ backgroundColor: 'text.secondary',
418
+ opacity: 0.8,
419
+ },
420
+ '& .MuiSlider-markActive': {
421
+ backgroundColor: 'text.secondary',
422
+ opacity: 0.8,
423
+ },
424
+ }}
425
+ />
426
+ </MidiMappable>
427
  </Box>
428
  </Box>
429
 
430
  <Box sx={styles.knobsGrid}>
431
+ {KNOB_DEFS.map((k) => {
432
+ const isLog = k.scale === 'log';
433
+ // For log knobs, the slider drives a 0..1 position and we
434
+ // convert to/from the underlying value (Hz) on the audio
435
+ // boundary. The knob value stored in state stays in the
436
+ // domain unit (Hz here) so persistence and MIDI keep working.
437
+ const valueToPos = isLog
438
+ ? (v) => Math.log(Math.max(v, k.min) / k.min) / Math.log(k.max / k.min)
439
+ : (v) => v;
440
+ const posToValue = isLog
441
+ ? (p) => k.min * Math.pow(k.max / k.min, p)
442
+ : (v) => v;
443
+ return (
444
+ <Box key={k.key} sx={styles.knobCell}>
445
+ <MidiMappable
446
+ id={ctrlId(k.key)}
447
+ label={ctrlLabel(k.label)}
448
+ kind="continuous"
449
+ curve={isLog ? 'log' : 'linear'}
450
+ min={k.min}
451
+ max={k.max}
452
+ value={knobs[k.key]}
453
+ onChange={(v) => handleKnob(k.key, v)}
454
+ sx={{ alignItems: 'center' }}
455
+ >
456
+ <Slider
457
+ orientation="vertical"
458
+ value={valueToPos(knobs[k.key])}
459
+ onChange={(_, v) => handleKnob(k.key, posToValue(v))}
460
+ min={isLog ? 0 : k.min}
461
+ max={isLog ? 1 : k.max}
462
+ step={isLog ? 0.001 : k.step}
463
+ size="small"
464
+ sx={styles.knobSlider(color, k.key === 'gain')}
465
+ />
466
+ </MidiMappable>
467
+ <Box component="span" sx={styles.knobLabel}>{k.label}</Box>
468
+ </Box>
469
+ );
470
+ })}
471
  </Box>
472
 
473
  <Box sx={styles.transportRow}>
474
+ <MidiMappable id={ctrlId('transport')} label={ctrlLabel('Play/Stop')} kind="trigger" onChange={handleTransportToggle}>
475
+ <IconButton
476
+ onClick={playing ? handleStop : handlePlay}
477
+ disabled={!loaded}
478
+ sx={styles.transportBtn(color, playing)}
479
+ size="small"
480
+ >
481
+ {playing ? <StopIcon size={16} /> : <PlayIcon size={16} />}
482
+ </IconButton>
483
+ </MidiMappable>
484
+ <MidiMappable id={ctrlId('loop')} label={ctrlLabel('Loop')} kind="trigger" onChange={handleLoopToggle}>
485
+ <IconButton
486
+ onClick={handleLoopToggle}
487
+ sx={styles.loopBtn(color, looping)}
488
+ size="small"
489
+ >
490
+ <LoopIcon size={14} />
491
+ </IconButton>
492
+ </MidiMappable>
493
  <Box sx={styles.meterTrack}>
494
  <Box ref={meterRef} sx={styles.meterFill(color)} />
495
  </Box>
app/frontend/src/components/PerformancePanel.js CHANGED
@@ -7,22 +7,39 @@ import {
7
  Button,
8
  Alert,
9
  FormControl,
 
 
10
  Select,
11
  MenuItem,
12
  TextField,
13
  IconButton,
14
  Tooltip,
15
  ButtonBase,
 
 
16
  } from '@mui/material';
17
  import {
18
  Play as PlayAllIcon,
19
  Square as StopAllIcon,
20
  Trash2 as DeleteIcon,
 
 
 
21
  } from 'lucide-react';
22
  import api from '../api';
23
  import PerformanceChannel from './PerformanceChannel';
24
  import { PerformanceEngine } from '../utils/performanceAudio';
25
- import { performancePanelStyles as styles } from '../theme';
 
 
 
 
 
 
 
 
 
 
26
 
27
  const CHANNEL_COUNT = 4;
28
  const MASTER_COLOR = '#35C2D4';
@@ -34,6 +51,21 @@ const BPM_MIN = 20;
34
  const BPM_MAX = 300;
35
  const BPM_DEFAULT = 120;
36
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37
  const dbToGain = (db) => (db <= MASTER_DB_MIN ? 0 : Math.pow(10, db / 20));
38
  const ampToDb = (amp) => (amp <= 0 ? -Infinity : 20 * Math.log10(amp));
39
  const formatDb = (db) => {
@@ -42,7 +74,15 @@ const formatDb = (db) => {
42
  return db.toFixed(1);
43
  };
44
 
45
- export default function PerformancePanel({
 
 
 
 
 
 
 
 
46
  selectedModel,
47
  selectedUnwrappedModel,
48
  availableModels = [],
@@ -50,32 +90,155 @@ export default function PerformancePanel({
50
  onSelectModel,
51
  onSelectUnwrappedModel,
52
  onRefreshModels,
 
 
 
 
 
 
 
53
  }) {
 
 
54
  const engineRef = useRef(null);
55
  const meterFillRef = useRef(null);
56
  const peakHoldRef = useRef({ db: METER_FLOOR_DB, decayedAt: performance.now() });
57
  const meterRafRef = useRef(null);
58
  const [engineReady, setEngineReady] = useState(false);
59
- const [masterDb, setMasterDb] = useState(MASTER_DB_DEFAULT);
60
- const [bpm, setBpm] = useState(BPM_DEFAULT);
61
- // Separate "what the field displays" from "what the app commits". Lets
62
- // the user type "1" → "11" → "112" without the first keystroke getting
63
- // clamped up to BPM_MIN mid-typing (which then produced "201", "211", etc).
64
- const [bpmInput, setBpmInput] = useState(String(BPM_DEFAULT));
65
  const bpmInputFocusedRef = useRef(false);
66
  const [error, setError] = useState(null);
67
  const [linkAvailable, setLinkAvailable] = useState(false);
68
- const [linkEnabled, setLinkEnabled] = useState(false);
69
  const [linkPeers, setLinkPeers] = useState(0);
70
  const [linkInstalling, setLinkInstalling] = useState(false);
71
- // Tracks whether the most recent bpm change came from Link (poll response)
72
- // or from the user (typing in the field). Prevents an echo loop where a
73
- // Link-driven update gets pushed back to Link as if it were a local edit.
74
  const bpmOriginRef = useRef('user');
75
  const [peakLabelDb, setPeakLabelDb] = useState(METER_FLOOR_DB);
76
  const [channelStates, setChannelStates] = useState(() =>
77
  Array.from({ length: CHANNEL_COUNT }, () => ({ loaded: false, playing: false }))
78
  );
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
79
 
80
  if (!engineRef.current) {
81
  engineRef.current = new PerformanceEngine(CHANNEL_COUNT);
@@ -129,9 +292,6 @@ export default function PerformancePanel({
129
  const handleBpmChange = (event) => {
130
  const raw = event.target.value;
131
  setBpmInput(raw);
132
- // Only commit the numeric bpm if what the user has typed so far is a
133
- // complete, in-range value. Intermediate digits ("1" on the way to
134
- // "112") are held in bpmInput without disturbing the committed bpm.
135
  const parsed = Number(raw);
136
  if (Number.isFinite(parsed) && parsed >= BPM_MIN && parsed <= BPM_MAX) {
137
  bpmOriginRef.current = 'user';
@@ -157,13 +317,10 @@ export default function PerformancePanel({
157
  bpmInputFocusedRef.current = true;
158
  };
159
 
160
- // Mirror committed bpm back into the field — but only when the user isn't
161
- // currently typing, so Link-driven updates don't overwrite a draft.
162
  useEffect(() => {
163
  if (!bpmInputFocusedRef.current) setBpmInput(String(bpm));
164
  }, [bpm]);
165
 
166
- // Probe whether the backend has an Ableton Link binding installed.
167
  useEffect(() => {
168
  api.get('/api/link/state')
169
  .then((r) => {
@@ -173,11 +330,16 @@ export default function PerformancePanel({
173
  .catch(() => setLinkAvailable(false));
174
  }, []);
175
 
176
- // Poll Link state while enabled; pull BPM and peer count into local state.
177
  useEffect(() => {
178
- if (!linkEnabled) return undefined;
 
 
 
 
179
  let cancelled = false;
180
  const poll = async () => {
 
181
  try {
182
  const r = await api.get('/api/link/state');
183
  if (cancelled || !r.data?.enabled) return;
@@ -190,6 +352,22 @@ export default function PerformancePanel({
190
  });
191
  }
192
  setLinkPeers(Number(r.data.num_peers || 0));
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
193
  } catch {
194
  /* transient network blip — next tick will retry */
195
  }
@@ -199,8 +377,10 @@ export default function PerformancePanel({
199
  return () => { cancelled = true; clearInterval(timer); };
200
  }, [linkEnabled]);
201
 
202
- // User-initiated BPM changes are pushed to the Link session. Changes that
203
- // originated from a Link poll are suppressed here (see bpmOriginRef).
 
 
204
  useEffect(() => {
205
  if (!linkEnabled) return;
206
  if (bpmOriginRef.current === 'link') {
@@ -222,7 +402,6 @@ export default function PerformancePanel({
222
  try {
223
  await api.post('/api/link/install');
224
  setLinkAvailable(true);
225
- // Auto-enable after a successful install — the user just asked for Link.
226
  await api.post('/api/link/enable');
227
  setLinkEnabled(true);
228
  } catch (err) {
@@ -251,23 +430,72 @@ export default function PerformancePanel({
251
  }
252
  }, [linkEnabled, linkAvailable]);
253
 
254
- const handlePlayAll = () => engineRef.current?.playAll(true);
255
- const handleStopAll = () => engineRef.current?.stopAll();
 
 
 
 
 
 
 
 
 
 
 
 
256
 
257
- const generateForChannel = async ({ prompt, duration }) => {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
258
  setError(null);
259
  if (!selectedModel) {
260
  const msg = 'Pick a model first.';
261
  setError(msg);
262
  throw new Error(msg);
263
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
264
  const requestData = {
265
- prompt,
266
  duration,
267
  cfg_scale: 7.0,
268
- seed: Math.floor(Math.random() * 0xffffffff),
 
269
  model_name: selectedModel,
270
  ...(selectedUnwrappedModel ? { unwrapped_model_path: selectedUnwrappedModel } : {}),
 
271
  };
272
  const response = await api.post('/api/generate', requestData, { responseType: 'blob' });
273
  return response.data;
@@ -370,18 +598,30 @@ export default function PerformancePanel({
370
  alignItems: 'center',
371
  justifyContent: 'center',
372
  fontFamily: 'inherit',
373
- fontSize: '0.72rem',
374
  fontWeight: 600,
375
  px: 1,
376
  minWidth: 46,
377
- height: 26,
378
  borderRadius: '2px',
379
- bgcolor: linkEnabled ? '#F5C542' : '#6e6e6e',
 
 
 
 
 
 
 
 
380
  color: linkEnabled ? '#000' : '#2a2a2a',
381
  opacity: linkInstalling ? 0.55 : 1,
382
  transition: 'background-color 120ms',
383
  '&:hover': {
384
- bgcolor: linkEnabled ? '#FFD54F' : '#7d7d7d',
 
 
 
 
385
  },
386
  '&.Mui-disabled': {
387
  color: linkEnabled ? '#000' : '#2a2a2a',
@@ -390,56 +630,321 @@ export default function PerformancePanel({
390
  >
391
  {linkInstalling
392
  ? 'installing…'
393
- : `Link${linkEnabled && linkPeers > 0 ? ` · ${linkPeers}` : ''}`}
 
 
394
  </ButtonBase>
395
  </span>
396
  </Tooltip>
397
 
398
- {/* Tempo — outlined field matching other inputs. The floating-label
399
- notch looked awkward at this width, so the unit label lives
400
- inline as an endAdornment instead. */}
401
- <TextField
402
- size="small"
403
- type="number"
404
- value={bpmInput}
405
- onChange={handleBpmChange}
406
- onFocus={handleBpmFocus}
407
- onBlur={handleBpmBlur}
408
- inputProps={{ step: 1, inputMode: 'numeric', 'aria-label': 'Tempo in BPM' }}
409
- InputProps={{
410
- endAdornment: (
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
411
  <Typography
412
- component="span"
413
- sx={{
414
- fontSize: '0.62rem',
415
- letterSpacing: '0.08em',
416
- color: 'text.disabled',
417
- pl: 0.5,
418
- userSelect: 'none',
419
- }}
420
  >
421
- BPM
422
  </Typography>
423
- ),
424
- }}
425
- sx={{
426
- width: 96,
427
- '& .MuiOutlinedInput-root': { borderRadius: 1.5, pr: 1 },
428
- '& input': {
429
- textAlign: 'right',
430
- fontFamily: 'ui-monospace, SFMono-Regular, Menlo, Consolas, monospace',
431
- fontVariantNumeric: 'tabular-nums',
432
- pr: 0,
433
- },
434
- '& input::-webkit-outer-spin-button, & input::-webkit-inner-spin-button': {
435
- WebkitAppearance: 'none',
436
- margin: 0,
437
- },
438
- '& input[type=number]': { MozAppearance: 'textfield' },
439
- }}
440
- />
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
441
 
442
- {/* Model picker — half the old width */}
443
  <FormControl size="small" sx={{
444
  flex: 1,
445
  minWidth: 110,
@@ -515,7 +1020,7 @@ export default function PerformancePanel({
515
  </Select>
516
  </FormControl>
517
 
518
- {/* Checkpoint picker — also halved */}
519
  {unwrappedModels.length > 0 && (
520
  <FormControl size="small" sx={{
521
  flex: 1,
@@ -548,27 +1053,30 @@ export default function PerformancePanel({
548
  </FormControl>
549
  )}
550
 
551
- {/* Transport */}
552
- <Button
553
- size="small"
554
- variant="outlined"
555
- startIcon={<PlayAllIcon size={14} />}
556
- onClick={handlePlayAll}
557
- disabled={!anyLoaded}
558
- sx={styles.masterBtn(MASTER_COLOR, 'play')}
559
- >
560
- Play All
561
- </Button>
562
- <Button
563
- size="small"
564
- variant="outlined"
565
- startIcon={<StopAllIcon size={14} />}
566
- onClick={handleStopAll}
567
- disabled={!anyPlaying}
568
- sx={styles.masterBtn(MASTER_COLOR, 'stop')}
569
- >
570
- Stop All
571
- </Button>
 
 
 
572
  </Paper>
573
 
574
  {error && (
@@ -584,10 +1092,14 @@ export default function PerformancePanel({
584
  key={i}
585
  index={i}
586
  strip={strip}
 
 
587
  onGenerate={generateForChannel}
588
  canGenerate={Boolean(selectedModel)}
589
  onMuteSoloChange={handleMuteSoloChange}
590
  onStateChange={handleChannelStateChange}
 
 
591
  maxDuration={maxDuration}
592
  bpm={bpm}
593
  />
@@ -604,15 +1116,26 @@ export default function PerformancePanel({
604
  <Box ref={meterFillRef} sx={styles.masterMeterFill(MASTER_COLOR)} />
605
  <Box sx={styles.masterMeterSegments} />
606
  </Box>
607
- <Slider
608
- orientation="vertical"
609
- value={masterDb}
610
- onChange={handleMasterChange}
611
  min={MASTER_DB_MIN}
612
  max={MASTER_DB_MAX}
613
- step={0.1}
614
- sx={styles.masterFader(MASTER_COLOR)}
615
- />
 
 
 
 
 
 
 
 
 
 
 
616
  </Box>
617
 
618
  <Box sx={styles.masterReadouts}>
@@ -625,6 +1148,106 @@ export default function PerformancePanel({
625
  </Box>
626
  </Box>
627
  </Box>
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
628
  </Box>
629
  );
630
  }
 
7
  Button,
8
  Alert,
9
  FormControl,
10
+ FormControlLabel,
11
+ Switch,
12
  Select,
13
  MenuItem,
14
  TextField,
15
  IconButton,
16
  Tooltip,
17
  ButtonBase,
18
+ Menu,
19
+ ListItemText,
20
  } from '@mui/material';
21
  import {
22
  Play as PlayAllIcon,
23
  Square as StopAllIcon,
24
  Trash2 as DeleteIcon,
25
+ Settings as SettingsIcon,
26
+ Save as SaveIcon,
27
+ X as CloseXIcon,
28
  } from 'lucide-react';
29
  import api from '../api';
30
  import PerformanceChannel from './PerformanceChannel';
31
  import { PerformanceEngine } from '../utils/performanceAudio';
32
+ import { performancePanelStyles as styles, perfTokens } from '../theme';
33
+ import { MidiProvider, MidiMappable, useMidi, clearMidiConfig } from './MidiContext';
34
+ import MidiConfigMenu from './MidiConfigMenu';
35
+ import {
36
+ usePerformanceSession,
37
+ listPresetNames,
38
+ savePreset,
39
+ deletePreset,
40
+ loadPresetIntoSession,
41
+ clearPerformanceSession,
42
+ } from './usePerformanceSession';
43
 
44
  const CHANNEL_COUNT = 4;
45
  const MASTER_COLOR = '#35C2D4';
 
51
  const BPM_MAX = 300;
52
  const BPM_DEFAULT = 120;
53
 
54
+
55
+ const LAUNCH_QUANTIZE_OPTIONS = [
56
+ { value: 0, label: 'None' },
57
+ { value: 32, label: '8 Bars' },
58
+ { value: 16, label: '4 Bars' },
59
+ { value: 8, label: '2 Bars' },
60
+ { value: 4, label: '1 Bar' },
61
+ { value: 2, label: '1/2' },
62
+ { value: 1, label: '1/4' },
63
+ { value: 0.5, label: '1/8' },
64
+ { value: 0.25, label: '1/16' },
65
+ { value: 0.125, label: '1/32' },
66
+ ];
67
+ const LAUNCH_Q_DEFAULT = 4;
68
+
69
  const dbToGain = (db) => (db <= MASTER_DB_MIN ? 0 : Math.pow(10, db / 20));
70
  const ampToDb = (amp) => (amp <= 0 ? -Infinity : 20 * Math.log10(amp));
71
  const formatDb = (db) => {
 
74
  return db.toFixed(1);
75
  };
76
 
77
+ export default function PerformancePanel(props) {
78
+ return (
79
+ <MidiProvider>
80
+ <PerformancePanelInner {...props} />
81
+ </MidiProvider>
82
+ );
83
+ }
84
+
85
+ function PerformancePanelInner({
86
  selectedModel,
87
  selectedUnwrappedModel,
88
  availableModels = [],
 
90
  onSelectModel,
91
  onSelectUnwrappedModel,
92
  onRefreshModels,
93
+ steps = 250,
94
+ onStepsChange,
95
+ randomSeed = true,
96
+ seedValue = '',
97
+ onRandomSeedChange,
98
+ onSeedValueChange,
99
+ onPresetLoaded,
100
  }) {
101
+ const { session, updateGlobal, updateChannel } = usePerformanceSession(CHANNEL_COUNT);
102
+
103
  const engineRef = useRef(null);
104
  const meterFillRef = useRef(null);
105
  const peakHoldRef = useRef({ db: METER_FLOOR_DB, decayedAt: performance.now() });
106
  const meterRafRef = useRef(null);
107
  const [engineReady, setEngineReady] = useState(false);
108
+ const [masterDb, setMasterDb] = useState(session.masterDb ?? MASTER_DB_DEFAULT);
109
+ const [bpm, setBpm] = useState(session.bpm ?? BPM_DEFAULT);
110
+ const [bpmInput, setBpmInput] = useState(String(session.bpm ?? BPM_DEFAULT));
 
 
 
111
  const bpmInputFocusedRef = useRef(false);
112
  const [error, setError] = useState(null);
113
  const [linkAvailable, setLinkAvailable] = useState(false);
114
+ const [linkEnabled, setLinkEnabled] = useState(session.linkEnabled ?? false);
115
  const [linkPeers, setLinkPeers] = useState(0);
116
  const [linkInstalling, setLinkInstalling] = useState(false);
117
+ const [launchQuantum, setLaunchQuantum] = useState(session.launchQuantum ?? LAUNCH_Q_DEFAULT);
118
+ const wasPlayingRef = useRef(false);
 
119
  const bpmOriginRef = useRef('user');
120
  const [peakLabelDb, setPeakLabelDb] = useState(METER_FLOOR_DB);
121
  const [channelStates, setChannelStates] = useState(() =>
122
  Array.from({ length: CHANNEL_COUNT }, () => ({ loaded: false, playing: false }))
123
  );
124
+ const [injectBpm, setInjectBpm] = useState(session.injectBpm ?? true);
125
+
126
+ // Restore App-level state (model, steps, seed) once on mount via the setter
127
+ // props the panel was given. The panel doesn't own those, so this is the
128
+ // only point at which we push session into App state. Subsequent changes
129
+ // flow normally through the prop callbacks.
130
+ const appStateRestoredRef = useRef(false);
131
+ useEffect(() => {
132
+ if (appStateRestoredRef.current) return;
133
+ appStateRestoredRef.current = true;
134
+ if (session.selectedModel && session.selectedModel !== selectedModel) {
135
+ onSelectModel?.(session.selectedModel);
136
+ }
137
+ if (session.selectedUnwrappedModel && session.selectedUnwrappedModel !== selectedUnwrappedModel) {
138
+ onSelectUnwrappedModel?.(session.selectedUnwrappedModel);
139
+ }
140
+ if (typeof session.steps === 'number' && session.steps !== steps) {
141
+ onStepsChange?.(session.steps);
142
+ }
143
+ if (typeof session.randomSeed === 'boolean' && session.randomSeed !== randomSeed) {
144
+ onRandomSeedChange?.(session.randomSeed);
145
+ }
146
+ if (typeof session.seedValue === 'string' && session.seedValue !== seedValue) {
147
+ onSeedValueChange?.(session.seedValue);
148
+ }
149
+ // Intentionally only run on first mount.
150
+ // eslint-disable-next-line react-hooks/exhaustive-deps
151
+ }, []);
152
+
153
+ // Push panel + App-level state into the session whenever any of it changes.
154
+ useEffect(() => { updateGlobal('bpm', bpm); }, [bpm, updateGlobal]);
155
+ useEffect(() => { updateGlobal('launchQuantum', launchQuantum); }, [launchQuantum, updateGlobal]);
156
+ useEffect(() => { updateGlobal('masterDb', masterDb); }, [masterDb, updateGlobal]);
157
+ useEffect(() => { updateGlobal('injectBpm', injectBpm); }, [injectBpm, updateGlobal]);
158
+ useEffect(() => { updateGlobal('linkEnabled', linkEnabled); }, [linkEnabled, updateGlobal]);
159
+ useEffect(() => { updateGlobal('selectedModel', selectedModel || ''); }, [selectedModel, updateGlobal]);
160
+ useEffect(() => { updateGlobal('selectedUnwrappedModel', selectedUnwrappedModel || ''); }, [selectedUnwrappedModel, updateGlobal]);
161
+ useEffect(() => { updateGlobal('steps', steps); }, [steps, updateGlobal]);
162
+ useEffect(() => { updateGlobal('randomSeed', randomSeed); }, [randomSeed, updateGlobal]);
163
+ useEffect(() => { updateGlobal('seedValue', seedValue); }, [seedValue, updateGlobal]);
164
+
165
+ const handleChannelFormChange = useCallback((index, partial) => {
166
+ updateChannel(index, partial);
167
+ }, [updateChannel]);
168
+
169
+ // ---- Preset menu state ----
170
+ const [presetMenuAnchor, setPresetMenuAnchor] = useState(null);
171
+ const [presetNames, setPresetNames] = useState(() => listPresetNames());
172
+ const [saveAsName, setSaveAsName] = useState('');
173
+ const [restoreArmed, setRestoreArmed] = useState(false);
174
+ const restoreArmTimerRef = useRef(null);
175
+
176
+ const refreshPresetNames = useCallback(() => {
177
+ setPresetNames(listPresetNames());
178
+ }, []);
179
+
180
+ const openPresetMenu = (e) => {
181
+ refreshPresetNames();
182
+ setSaveAsName('');
183
+ setRestoreArmed(false);
184
+ setPresetMenuAnchor(e.currentTarget);
185
+ };
186
+ const closePresetMenu = () => {
187
+ setPresetMenuAnchor(null);
188
+ setRestoreArmed(false);
189
+ if (restoreArmTimerRef.current) {
190
+ clearTimeout(restoreArmTimerRef.current);
191
+ restoreArmTimerRef.current = null;
192
+ }
193
+ };
194
+
195
+ const handleRestoreDefaults = () => {
196
+ if (!restoreArmed) {
197
+ // First click arms; second click within 3 s commits. Disarms
198
+ // automatically so the destructive path is never one accidental
199
+ // click away.
200
+ setRestoreArmed(true);
201
+ if (restoreArmTimerRef.current) clearTimeout(restoreArmTimerRef.current);
202
+ restoreArmTimerRef.current = setTimeout(() => setRestoreArmed(false), 3000);
203
+ return;
204
+ }
205
+ clearPerformanceSession();
206
+ clearMidiConfig();
207
+ closePresetMenu();
208
+ onPresetLoaded?.();
209
+ };
210
+
211
+ const handleSaveAs = () => {
212
+ const name = saveAsName.trim();
213
+ if (!name) return;
214
+ savePreset(name, session);
215
+ setSaveAsName('');
216
+ refreshPresetNames();
217
+ };
218
+
219
+ const handleLoadPreset = (name) => {
220
+ if (!loadPresetIntoSession(name)) return;
221
+ closePresetMenu();
222
+ // Force-remount via the App-level reset key. Same pathway as Fresh
223
+ // Start, just with a different localStorage payload pre-loaded.
224
+ onPresetLoaded?.();
225
+ };
226
+
227
+ const handleDeletePreset = (name, e) => {
228
+ e?.stopPropagation();
229
+ deletePreset(name);
230
+ refreshPresetNames();
231
+ };
232
+
233
+ const isSmallModel = (() => {
234
+ if (selectedModel === 'stable-audio-open-small') return true;
235
+ const model = availableModels.find((m) => m.name === selectedModel);
236
+ if (model && selectedUnwrappedModel) {
237
+ const u = model.unwrapped_models?.find((x) => x.path === selectedUnwrappedModel);
238
+ return u ? (u.size_mb || 0) < 2000 : false;
239
+ }
240
+ return false;
241
+ })();
242
 
243
  if (!engineRef.current) {
244
  engineRef.current = new PerformanceEngine(CHANNEL_COUNT);
 
292
  const handleBpmChange = (event) => {
293
  const raw = event.target.value;
294
  setBpmInput(raw);
 
 
 
295
  const parsed = Number(raw);
296
  if (Number.isFinite(parsed) && parsed >= BPM_MIN && parsed <= BPM_MAX) {
297
  bpmOriginRef.current = 'user';
 
317
  bpmInputFocusedRef.current = true;
318
  };
319
 
 
 
320
  useEffect(() => {
321
  if (!bpmInputFocusedRef.current) setBpmInput(String(bpm));
322
  }, [bpm]);
323
 
 
324
  useEffect(() => {
325
  api.get('/api/link/state')
326
  .then((r) => {
 
330
  .catch(() => setLinkAvailable(false));
331
  }, []);
332
 
333
+
334
  useEffect(() => {
335
+ if (!linkEnabled) {
336
+ engineRef.current?.setLinkSnapshot(null);
337
+ wasPlayingRef.current = false;
338
+ return undefined;
339
+ }
340
  let cancelled = false;
341
  const poll = async () => {
342
+ const capturedAt = performance.now();
343
  try {
344
  const r = await api.get('/api/link/state');
345
  if (cancelled || !r.data?.enabled) return;
 
352
  });
353
  }
354
  setLinkPeers(Number(r.data.num_peers || 0));
355
+
356
+ const isPlaying = Boolean(r.data.is_playing);
357
+ const beat = Number(r.data.beat) || 0;
358
+ const bpmFloat = Number(r.data.bpm) || 120;
359
+ engineRef.current?.setLinkSnapshot({
360
+ beat,
361
+ bpm: bpmFloat,
362
+ isPlaying,
363
+ capturedAt,
364
+ });
365
+
366
+ if (wasPlayingRef.current && !isPlaying) {
367
+ engineRef.current?.stopAll();
368
+ setChannelStates(prev => prev.map(s => ({ ...s, playing: false })));
369
+ }
370
+ wasPlayingRef.current = isPlaying;
371
  } catch {
372
  /* transient network blip — next tick will retry */
373
  }
 
377
  return () => { cancelled = true; clearInterval(timer); };
378
  }, [linkEnabled]);
379
 
380
+ useEffect(() => {
381
+ engineRef.current?.setLaunchQuantum(launchQuantum);
382
+ }, [launchQuantum]);
383
+
384
  useEffect(() => {
385
  if (!linkEnabled) return;
386
  if (bpmOriginRef.current === 'link') {
 
402
  try {
403
  await api.post('/api/link/install');
404
  setLinkAvailable(true);
 
405
  await api.post('/api/link/enable');
406
  setLinkEnabled(true);
407
  } catch (err) {
 
430
  }
431
  }, [linkEnabled, linkAvailable]);
432
 
433
+ const handlePlayAll = () => {
434
+ engineRef.current?.playAll(true);
435
+ setChannelStates(prev => prev.map(s => (s.loaded ? { ...s, playing: true } : s)));
436
+ };
437
+ const handleStopAll = () => {
438
+ engineRef.current?.stopAll();
439
+ setChannelStates(prev => prev.map(s => ({ ...s, playing: false })));
440
+ };
441
+
442
+ const applyExternalBpm = useCallback((value) => {
443
+ const next = Math.max(BPM_MIN, Math.min(BPM_MAX, Math.round(value)));
444
+ bpmOriginRef.current = 'user';
445
+ setBpm(next);
446
+ }, []);
447
 
448
+ const midi = useMidi();
449
+ const [midiMenuAnchor, setMidiMenuAnchor] = useState(null);
450
+
451
+ useEffect(() => {
452
+ if (!midi?.learnMode) return undefined;
453
+ const onKey = (e) => {
454
+ if (e.key === 'Escape') {
455
+ e.preventDefault();
456
+ midi.exitLearnMode();
457
+ }
458
+ };
459
+ window.addEventListener('keydown', onKey);
460
+ return () => window.removeEventListener('keydown', onKey);
461
+ }, [midi?.learnMode, midi?.exitLearnMode]);
462
+
463
+ const generateForChannel = async ({ prompt, duration, alignBars, alignBpm }) => {
464
  setError(null);
465
  if (!selectedModel) {
466
  const msg = 'Pick a model first.';
467
  setError(msg);
468
  throw new Error(msg);
469
  }
470
+
471
+ const trimmed = (prompt || '').trim();
472
+ const hasExplicitBpm = /\b\d{2,3}\s*bpm\b/i.test(trimmed);
473
+ const finalPrompt = injectBpm && !hasExplicitBpm
474
+ ? `${trimmed}${trimmed ? ', ' : ''}${Math.round(bpm)} BPM`
475
+ : trimmed;
476
+
477
+ let resolvedSeed;
478
+ if (randomSeed) {
479
+ resolvedSeed = Math.floor(Math.random() * 0xffffffff);
480
+ } else {
481
+ const parsed = parseInt(seedValue, 10);
482
+ if (Number.isNaN(parsed) || parsed < 0) {
483
+ const msg = 'Enter a valid seed (0 or greater) or enable Random.';
484
+ setError(msg);
485
+ throw new Error(msg);
486
+ }
487
+ resolvedSeed = parsed;
488
+ }
489
+
490
  const requestData = {
491
+ prompt: finalPrompt,
492
  duration,
493
  cfg_scale: 7.0,
494
+ steps,
495
+ seed: resolvedSeed,
496
  model_name: selectedModel,
497
  ...(selectedUnwrappedModel ? { unwrapped_model_path: selectedUnwrappedModel } : {}),
498
+ ...(alignBars && alignBpm ? { align_bars: alignBars, align_bpm: alignBpm } : {}),
499
  };
500
  const response = await api.post('/api/generate', requestData, { responseType: 'blob' });
501
  return response.data;
 
598
  alignItems: 'center',
599
  justifyContent: 'center',
600
  fontFamily: 'inherit',
601
+ fontSize: perfTokens.fontSize.body,
602
  fontWeight: 600,
603
  px: 1,
604
  minWidth: 46,
605
+ height: perfTokens.height.compact,
606
  borderRadius: '2px',
607
+ // Three states matching Ableton Live's button:
608
+ // off → gray, on/no peers → yellow (broadcasting,
609
+ // alone), on/peers → teal (sync'd with at least
610
+ // one other app).
611
+ bgcolor: !linkEnabled
612
+ ? '#6e6e6e'
613
+ : linkPeers > 0
614
+ ? MASTER_COLOR
615
+ : '#F5C542',
616
  color: linkEnabled ? '#000' : '#2a2a2a',
617
  opacity: linkInstalling ? 0.55 : 1,
618
  transition: 'background-color 120ms',
619
  '&:hover': {
620
+ bgcolor: !linkEnabled
621
+ ? '#7d7d7d'
622
+ : linkPeers > 0
623
+ ? '#4DD0DE'
624
+ : '#FFD54F',
625
  },
626
  '&.Mui-disabled': {
627
  color: linkEnabled ? '#000' : '#2a2a2a',
 
630
  >
631
  {linkInstalling
632
  ? 'installing…'
633
+ : linkEnabled && linkPeers > 0
634
+ ? `${linkPeers} Link`
635
+ : 'Link'}
636
  </ButtonBase>
637
  </span>
638
  </Tooltip>
639
 
640
+ {/* MIDI learn toggle — same compact rectangle style as Link. */}
641
+ <Tooltip
642
+ title={
643
+ !midi?.supported
644
+ ? (midi?.permissionError || 'Web MIDI is not available')
645
+ : midi.learnMode
646
+ ? 'Exit MIDI mode (Esc)'
647
+ : 'Enter MIDI mode — click a control then move a hardware knob/button to bind'
648
+ }
649
+ >
650
+ <span style={{ display: 'inline-flex', alignItems: 'center' }}>
651
+ <ButtonBase
652
+ onClick={() => midi?.toggleLearnMode()}
653
+ disabled={!midi?.supported}
654
+ sx={{
655
+ display: 'inline-flex',
656
+ alignItems: 'center',
657
+ justifyContent: 'center',
658
+ fontFamily: 'inherit',
659
+ fontSize: perfTokens.fontSize.body,
660
+ fontWeight: 600,
661
+ px: 1,
662
+ minWidth: 46,
663
+ height: perfTokens.height.compact,
664
+ borderRadius: '2px',
665
+ bgcolor: midi?.learnMode ? '#F5C542' : '#6e6e6e',
666
+ color: midi?.learnMode ? '#000' : '#2a2a2a',
667
+ opacity: midi?.supported ? 1 : 0.45,
668
+ transition: 'background-color 120ms',
669
+ '&:hover': {
670
+ bgcolor: midi?.learnMode ? '#FFD54F' : '#7d7d7d',
671
+ },
672
+ '&.Mui-disabled': {
673
+ color: '#2a2a2a',
674
+ },
675
+ }}
676
+ >
677
+ MIDI
678
+ </ButtonBase>
679
+ </span>
680
+ </Tooltip>
681
+ <Tooltip title="MIDI settings & mappings">
682
+ <span style={{ display: 'inline-flex', alignItems: 'center' }}>
683
+ <IconButton
684
+ size="small"
685
+ onClick={(e) => setMidiMenuAnchor(e.currentTarget)}
686
+ sx={{ width: perfTokens.height.compact, height: perfTokens.height.compact, color: 'text.secondary' }}
687
+ >
688
+ <SettingsIcon size={14} />
689
+ </IconButton>
690
+ </span>
691
+ </Tooltip>
692
+ <MidiConfigMenu
693
+ anchorEl={midiMenuAnchor}
694
+ open={Boolean(midiMenuAnchor)}
695
+ onClose={() => setMidiMenuAnchor(null)}
696
+ />
697
+
698
+ <Tooltip title="Save / load presets">
699
+ <span style={{ display: 'inline-flex', alignItems: 'center' }}>
700
+ <IconButton
701
+ size="small"
702
+ onClick={openPresetMenu}
703
+ sx={{ width: perfTokens.height.compact, height: perfTokens.height.compact, color: 'text.secondary' }}
704
+ >
705
+ <SaveIcon size={14} />
706
+ </IconButton>
707
+ </span>
708
+ </Tooltip>
709
+ <Menu
710
+ anchorEl={presetMenuAnchor}
711
+ open={Boolean(presetMenuAnchor)}
712
+ onClose={closePresetMenu}
713
+ MenuListProps={{ sx: { py: 0 } }}
714
+ PaperProps={{
715
+ sx: {
716
+ minWidth: 240,
717
+ borderRadius: 1.5,
718
+ },
719
+ }}
720
+ >
721
+ {/* SAVE section — type a name and click save (or press Enter) */}
722
+ <Box sx={{ px: 1.5, pt: 1.25, pb: 0.5 }}>
723
+ <Typography
724
+ sx={{
725
+ fontSize: perfTokens.fontSize.badge,
726
+ letterSpacing: perfTokens.letterSpacing.wide,
727
+ fontWeight: 700,
728
+ color: 'text.secondary',
729
+ textTransform: 'uppercase',
730
+ }}
731
+ >
732
+ Save
733
+ </Typography>
734
+ </Box>
735
+ <Box sx={{ px: 1.5, pb: 1.25, display: 'flex', alignItems: 'center', gap: 0.75 }}>
736
+ <TextField
737
+ autoFocus
738
+ size="small"
739
+ placeholder="Preset name"
740
+ value={saveAsName}
741
+ onChange={(e) => setSaveAsName(e.target.value)}
742
+ onKeyDown={(e) => {
743
+ if (e.key === 'Enter') handleSaveAs();
744
+ e.stopPropagation();
745
+ }}
746
+ sx={{
747
+ flex: 1,
748
+ '& .MuiOutlinedInput-root': {
749
+ borderRadius: 1.5,
750
+ height: perfTokens.height.compact,
751
+ fontSize: perfTokens.fontSize.body,
752
+ },
753
+ }}
754
+ />
755
+ <Button
756
+ size="small"
757
+ variant="contained"
758
+ onClick={handleSaveAs}
759
+ disabled={!saveAsName.trim()}
760
+ sx={{
761
+ height: perfTokens.height.compact,
762
+ fontSize: perfTokens.fontSize.body,
763
+ fontWeight: 600,
764
+ letterSpacing: perfTokens.letterSpacing.wide,
765
+ borderRadius: 1.5,
766
+ minWidth: 56,
767
+ px: 1.25,
768
+ }}
769
+ >
770
+ Save
771
+ </Button>
772
+ </Box>
773
+ {saveAsName.trim() && presetNames.includes(saveAsName.trim()) && (
774
+ <Box sx={{ px: 1.5, pb: 0.75 }}>
775
  <Typography
776
+ variant="caption"
777
+ color="warning.main"
778
+ sx={{ fontSize: perfTokens.fontSize.small }}
 
 
 
 
 
779
  >
780
+ Will overwrite existing preset.
781
  </Typography>
782
+ </Box>
783
+ )}
784
+
785
+ <Box sx={{ borderTop: '1px solid', borderColor: 'divider' }} />
786
+
787
+ {/* LOAD section — list of saved presets */}
788
+ <Box sx={{ px: 1.5, pt: 1.25, pb: 0.5 }}>
789
+ <Typography
790
+ sx={{
791
+ fontSize: perfTokens.fontSize.badge,
792
+ letterSpacing: perfTokens.letterSpacing.wide,
793
+ fontWeight: 700,
794
+ color: 'text.secondary',
795
+ textTransform: 'uppercase',
796
+ }}
797
+ >
798
+ Load
799
+ </Typography>
800
+ </Box>
801
+ {presetNames.length === 0 ? (
802
+ <Box sx={{ px: 1.5, pb: 1.25 }}>
803
+ <Typography
804
+ variant="caption"
805
+ color="text.disabled"
806
+ sx={{ fontSize: perfTokens.fontSize.small }}
807
+ >
808
+ No presets saved yet.
809
+ </Typography>
810
+ </Box>
811
+ ) : (
812
+ <Box sx={{ pb: 0.5 }}>
813
+ {presetNames.map((name) => (
814
+ <MenuItem
815
+ key={name}
816
+ onClick={() => handleLoadPreset(name)}
817
+ sx={{
818
+ display: 'flex',
819
+ justifyContent: 'space-between',
820
+ gap: 1,
821
+ fontSize: perfTokens.fontSize.body,
822
+ fontWeight: 600,
823
+ letterSpacing: perfTokens.letterSpacing.wide,
824
+ }}
825
+ >
826
+ <ListItemText primary={name} primaryTypographyProps={{ sx: { fontSize: perfTokens.fontSize.body } }} />
827
+ <Tooltip title="Delete preset">
828
+ <IconButton
829
+ size="small"
830
+ onClick={(e) => handleDeletePreset(name, e)}
831
+ sx={{ ml: 1, width: perfTokens.height.sub, height: perfTokens.height.sub }}
832
+ >
833
+ <CloseXIcon size={12} />
834
+ </IconButton>
835
+ </Tooltip>
836
+ </MenuItem>
837
+ ))}
838
+ </Box>
839
+ )}
840
+
841
+ <Box sx={{ borderTop: '1px solid', borderColor: 'divider' }} />
842
+
843
+ {/* Destructive: wipes session + MIDI mappings. Two-click arm
844
+ prevents accidental clicks; the menu auto-disarms after 3 s. */}
845
+ <Tooltip
846
+ title={restoreArmed
847
+ ? 'Click again within 3s to confirm — this clears all panel settings AND MIDI mappings'
848
+ : 'Reset panel settings and clear MIDI mappings'}
849
+ >
850
+ <MenuItem
851
+ onClick={handleRestoreDefaults}
852
+ sx={{
853
+ fontSize: perfTokens.fontSize.body,
854
+ fontWeight: 600,
855
+ letterSpacing: perfTokens.letterSpacing.wide,
856
+ color: restoreArmed ? 'error.main' : 'text.secondary',
857
+ }}
858
+ >
859
+ <ListItemText
860
+ primary={restoreArmed ? 'Click again to confirm' : 'Restore defaults'}
861
+ primaryTypographyProps={{ sx: { fontSize: perfTokens.fontSize.body } }}
862
+ />
863
+ </MenuItem>
864
+ </Tooltip>
865
+ </Menu>
866
+
867
+ <Tooltip placement="right" title="Launch quantization — match Live's">
868
+ <FormControl
869
+ size="small"
870
+ sx={{
871
+ minWidth: 92,
872
+ '& .MuiOutlinedInput-root': { borderRadius: 1.5, height: perfTokens.height.compact },
873
+ '& .MuiSelect-select': {
874
+ py: 0,
875
+ fontSize: perfTokens.fontSize.body,
876
+ fontWeight: 600,
877
+ letterSpacing: perfTokens.letterSpacing.wide,
878
+ },
879
+ }}
880
+ >
881
+ <Select
882
+ value={launchQuantum}
883
+ onChange={(e) => setLaunchQuantum(Number(e.target.value))}
884
+ renderValue={(val) => {
885
+ const opt = LAUNCH_QUANTIZE_OPTIONS.find((o) => o.value === val);
886
+ return `Q · ${opt?.label ?? 'None'}`;
887
+ }}
888
+ >
889
+ {LAUNCH_QUANTIZE_OPTIONS.map((opt) => (
890
+ <MenuItem key={opt.value} value={opt.value}>
891
+ <Typography variant="body2">{opt.label}</Typography>
892
+ </MenuItem>
893
+ ))}
894
+ </Select>
895
+ </FormControl>
896
+ </Tooltip>
897
+
898
+ <MidiMappable
899
+ id="master.bpm"
900
+ label="Tempo (BPM)"
901
+ kind="continuous"
902
+ min={BPM_MIN}
903
+ max={BPM_MAX}
904
+ value={bpm}
905
+ onChange={applyExternalBpm}
906
+ >
907
+ <TextField
908
+ size="small"
909
+ type="number"
910
+ value={bpmInput}
911
+ onChange={handleBpmChange}
912
+ onFocus={handleBpmFocus}
913
+ onBlur={handleBpmBlur}
914
+ inputProps={{ step: 1, inputMode: 'numeric', 'aria-label': 'Tempo in BPM' }}
915
+ InputProps={{
916
+ endAdornment: (
917
+ <Typography
918
+ component="span"
919
+ sx={{
920
+ fontSize: perfTokens.fontSize.small,
921
+ letterSpacing: perfTokens.letterSpacing.wide,
922
+ color: 'text.disabled',
923
+ pl: 0.5,
924
+ userSelect: 'none',
925
+ }}
926
+ >
927
+ BPM
928
+ </Typography>
929
+ ),
930
+ }}
931
+ sx={{
932
+ width: 96,
933
+ '& .MuiOutlinedInput-root': { borderRadius: 1.5, pr: 1 },
934
+ '& input': {
935
+ textAlign: 'right',
936
+ fontVariantNumeric: 'tabular-nums',
937
+ pr: 0,
938
+ },
939
+ '& input::-webkit-outer-spin-button, & input::-webkit-inner-spin-button': {
940
+ WebkitAppearance: 'none',
941
+ margin: 0,
942
+ },
943
+ '& input[type=number]': { MozAppearance: 'textfield' },
944
+ }}
945
+ />
946
+ </MidiMappable>
947
 
 
948
  <FormControl size="small" sx={{
949
  flex: 1,
950
  minWidth: 110,
 
1020
  </Select>
1021
  </FormControl>
1022
 
1023
+
1024
  {unwrappedModels.length > 0 && (
1025
  <FormControl size="small" sx={{
1026
  flex: 1,
 
1053
  </FormControl>
1054
  )}
1055
 
1056
+ <MidiMappable id="master.playAll" label="Play All" kind="trigger" onChange={handlePlayAll}>
1057
+ <Button
1058
+ size="small"
1059
+ variant="outlined"
1060
+ startIcon={<PlayAllIcon size={14} />}
1061
+ onClick={handlePlayAll}
1062
+ disabled={!anyLoaded}
1063
+ sx={styles.masterBtn(MASTER_COLOR, 'play')}
1064
+ >
1065
+ Play All
1066
+ </Button>
1067
+ </MidiMappable>
1068
+ <MidiMappable id="master.stopAll" label="Stop All" kind="trigger" onChange={handleStopAll}>
1069
+ <Button
1070
+ size="small"
1071
+ variant="outlined"
1072
+ startIcon={<StopAllIcon size={14} />}
1073
+ onClick={handleStopAll}
1074
+ disabled={!anyPlaying}
1075
+ sx={styles.masterBtn(MASTER_COLOR, 'stop')}
1076
+ >
1077
+ Stop All
1078
+ </Button>
1079
+ </MidiMappable>
1080
  </Paper>
1081
 
1082
  {error && (
 
1092
  key={i}
1093
  index={i}
1094
  strip={strip}
1095
+ engine={engineRef.current}
1096
+ playing={channelStates[i]?.playing || false}
1097
  onGenerate={generateForChannel}
1098
  canGenerate={Boolean(selectedModel)}
1099
  onMuteSoloChange={handleMuteSoloChange}
1100
  onStateChange={handleChannelStateChange}
1101
+ onFormStateChange={handleChannelFormChange}
1102
+ initialFormState={session.channels[i]}
1103
  maxDuration={maxDuration}
1104
  bpm={bpm}
1105
  />
 
1116
  <Box ref={meterFillRef} sx={styles.masterMeterFill(MASTER_COLOR)} />
1117
  <Box sx={styles.masterMeterSegments} />
1118
  </Box>
1119
+ <MidiMappable
1120
+ id="master.fader"
1121
+ label="Master Fader"
1122
+ kind="continuous"
1123
  min={MASTER_DB_MIN}
1124
  max={MASTER_DB_MAX}
1125
+ value={masterDb}
1126
+ onChange={(v) => handleMasterChange(null, v)}
1127
+ sx={{ flex: 1, alignSelf: 'stretch' }}
1128
+ >
1129
+ <Slider
1130
+ orientation="vertical"
1131
+ value={masterDb}
1132
+ onChange={handleMasterChange}
1133
+ min={MASTER_DB_MIN}
1134
+ max={MASTER_DB_MAX}
1135
+ step={0.1}
1136
+ sx={styles.masterFader(MASTER_COLOR)}
1137
+ />
1138
+ </MidiMappable>
1139
  </Box>
1140
 
1141
  <Box sx={styles.masterReadouts}>
 
1148
  </Box>
1149
  </Box>
1150
  </Box>
1151
+
1152
+ <Paper sx={{
1153
+ display: 'flex',
1154
+ alignItems: 'center',
1155
+ gap: 2.5,
1156
+ px: 1.5,
1157
+ py: 0.75,
1158
+ mt: 1,
1159
+ borderRadius: 2,
1160
+ border: '1px solid',
1161
+ borderColor: 'divider',
1162
+ background: 'linear-gradient(135deg, rgba(53, 194, 212, 0.04) 0%, rgba(159, 138, 230, 0.03) 100%)',
1163
+ flexWrap: { xs: 'wrap', md: 'nowrap' },
1164
+ }}>
1165
+ <Box sx={{ display: 'flex', alignItems: 'center', gap: 1 }}>
1166
+ <Typography variant="caption" color="textSecondary" sx={{ letterSpacing: perfTokens.letterSpacing.wide }}>
1167
+ STEPS
1168
+ </Typography>
1169
+ <Tooltip
1170
+ placement="right"
1171
+ title={
1172
+ isSmallModel
1173
+ ? 'Locked at 8 steps for the distilled small model'
1174
+ : 'Diffusion steps per generation (more = higher quality, slower)'
1175
+ }
1176
+ >
1177
+ <FormControl
1178
+ size="small"
1179
+ sx={{ minWidth: 96, '& .MuiOutlinedInput-root': { borderRadius: 1.5 } }}
1180
+ >
1181
+ <Select
1182
+ value={isSmallModel ? 8 : steps}
1183
+ onChange={(e) => onStepsChange?.(Number(e.target.value))}
1184
+ disabled={isSmallModel}
1185
+ renderValue={(value) => `${value} steps`}
1186
+ >
1187
+ {isSmallModel && (
1188
+ <MenuItem value={8}>
1189
+ <Typography variant="body2">8 (locked)</Typography>
1190
+ </MenuItem>
1191
+ )}
1192
+ {[50, 100, 150, 200, 250].map((n) => (
1193
+ <MenuItem key={n} value={n}>
1194
+ <Typography variant="body2">{n} steps</Typography>
1195
+ </MenuItem>
1196
+ ))}
1197
+ </Select>
1198
+ </FormControl>
1199
+ </Tooltip>
1200
+ </Box>
1201
+
1202
+ <Box sx={{ display: 'flex', alignItems: 'center', gap: 1 }}>
1203
+ <Typography variant="caption" color="textSecondary" sx={{ letterSpacing: perfTokens.letterSpacing.wide }}>
1204
+ SEED
1205
+ </Typography>
1206
+ <FormControlLabel
1207
+ sx={{ mr: 0 }}
1208
+ control={
1209
+ <Switch
1210
+ size="small"
1211
+ checked={randomSeed}
1212
+ onChange={(e) => onRandomSeedChange?.(e.target.checked)}
1213
+ />
1214
+ }
1215
+ label={<Typography variant="caption">Random</Typography>}
1216
+ />
1217
+ <TextField
1218
+ size="small"
1219
+ type="number"
1220
+ placeholder="e.g. 42"
1221
+ value={seedValue}
1222
+ onChange={(e) => onSeedValueChange?.(e.target.value)}
1223
+ disabled={randomSeed}
1224
+ inputProps={{ min: 0, max: 4294967295, step: 1 }}
1225
+ sx={{
1226
+ width: 130,
1227
+ '& .MuiOutlinedInput-root': { borderRadius: 1.5 },
1228
+ '& input': {
1229
+ fontVariantNumeric: 'tabular-nums',
1230
+ },
1231
+ }}
1232
+ />
1233
+ </Box>
1234
+
1235
+ <Tooltip
1236
+ placement="right"
1237
+ title="When on, the master BPM is injected to each prompt automatically (turn off if doing free-tempo or multi-tempo prompts)."
1238
+ >
1239
+ <Box sx={{ display: 'flex', alignItems: 'center', gap: 1 }}>
1240
+ <Typography variant="caption" color="textSecondary" sx={{ letterSpacing: perfTokens.letterSpacing.wide }}>
1241
+ AUTO BPM
1242
+ </Typography>
1243
+ <Switch
1244
+ size="small"
1245
+ checked={injectBpm}
1246
+ onChange={(e) => setInjectBpm(e.target.checked)}
1247
+ />
1248
+ </Box>
1249
+ </Tooltip>
1250
+ </Paper>
1251
  </Box>
1252
  );
1253
  }
app/frontend/src/components/TabPanel.js CHANGED
@@ -2,16 +2,37 @@ import React from 'react';
2
  import { Box } from '@mui/material';
3
  import { tabPanelStyles } from '../theme';
4
 
5
- export default function TabPanel({ children, value, index, ...other }) {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6
  return (
7
  <div
8
  role="tabpanel"
9
- hidden={value !== index}
10
  id={`simple-tabpanel-${index}`}
11
  aria-labelledby={`simple-tab-${index}`}
12
  {...other}
13
  >
14
- {value === index && (
15
  <Box sx={tabPanelStyles.root}>
16
  {children}
17
  </Box>
 
2
  import { Box } from '@mui/material';
3
  import { tabPanelStyles } from '../theme';
4
 
5
+ export default function TabPanel({ children, value, index, keepMounted = false, ...other }) {
6
+ const isActive = value === index;
7
+ // keepMounted: render children unconditionally and toggle visibility via CSS,
8
+ // so component state, audio nodes, and decoded buffers survive tab switches.
9
+ // Use sparingly — by default we still mount/unmount so inactive tabs cost
10
+ // nothing at idle.
11
+ if (keepMounted) {
12
+ return (
13
+ <div
14
+ role="tabpanel"
15
+ hidden={!isActive}
16
+ id={`simple-tabpanel-${index}`}
17
+ aria-labelledby={`simple-tab-${index}`}
18
+ {...other}
19
+ >
20
+ <Box sx={{ ...tabPanelStyles.root, display: isActive ? undefined : 'none' }}>
21
+ {children}
22
+ </Box>
23
+ </div>
24
+ );
25
+ }
26
+
27
  return (
28
  <div
29
  role="tabpanel"
30
+ hidden={!isActive}
31
  id={`simple-tabpanel-${index}`}
32
  aria-labelledby={`simple-tab-${index}`}
33
  {...other}
34
  >
35
+ {isActive && (
36
  <Box sx={tabPanelStyles.root}>
37
  {children}
38
  </Box>
app/frontend/src/components/WelcomePage.js CHANGED
@@ -67,7 +67,7 @@ export default function WelcomePage({ open, onClose }) {
67
  variant="body2"
68
  sx={welcomePageStyles.version}
69
  >
70
- Version 0.0.2
71
  </Typography>
72
  <Button
73
  variant="contained"
 
67
  variant="body2"
68
  sx={welcomePageStyles.version}
69
  >
70
+ Version 0.1.0
71
  </Typography>
72
  <Button
73
  variant="contained"
app/frontend/src/components/usePerformanceSession.js ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useCallback, useEffect, useRef, useState } from 'react';
2
+
3
+ export const PERFORMANCE_SESSION_STORAGE_KEY = 'fragmenta.performance.session.v1';
4
+
5
+ const STORAGE_KEY = PERFORMANCE_SESSION_STORAGE_KEY;
6
+
7
+ export function clearPerformanceSession() {
8
+ try { localStorage.removeItem(STORAGE_KEY); }
9
+ catch { /* non-fatal */ }
10
+ }
11
+
12
+ const PRESETS_STORAGE_KEY = 'fragmenta.performance.presets.v1';
13
+
14
+ function readPresetBag() {
15
+ try {
16
+ const raw = localStorage.getItem(PRESETS_STORAGE_KEY);
17
+ if (!raw) return {};
18
+ const parsed = JSON.parse(raw);
19
+ return parsed && typeof parsed === 'object' ? parsed : {};
20
+ } catch {
21
+ return {};
22
+ }
23
+ }
24
+
25
+ function writePresetBag(bag) {
26
+ try { localStorage.setItem(PRESETS_STORAGE_KEY, JSON.stringify(bag)); }
27
+ catch { /* quota — non-fatal */ }
28
+ }
29
+
30
+ export function listPresetNames() {
31
+ return Object.keys(readPresetBag()).sort((a, b) => a.localeCompare(b));
32
+ }
33
+
34
+ export function savePreset(name, sessionData) {
35
+ const trimmed = (name || '').trim();
36
+ if (!trimmed) return false;
37
+ const bag = readPresetBag();
38
+ bag[trimmed] = sessionData;
39
+ writePresetBag(bag);
40
+ return true;
41
+ }
42
+
43
+ export function deletePreset(name) {
44
+ const bag = readPresetBag();
45
+ if (!(name in bag)) return false;
46
+ delete bag[name];
47
+ writePresetBag(bag);
48
+ return true;
49
+ }
50
+
51
+ // Replace the live session storage with a preset's snapshot. Caller is
52
+ // expected to force-remount the panel afterward so its useState mirrors
53
+ // pick up the new shape; localStorage alone won't reset mounted state.
54
+ export function loadPresetIntoSession(name) {
55
+ const bag = readPresetBag();
56
+ const preset = bag[name];
57
+ if (!preset) return false;
58
+ try {
59
+ localStorage.setItem(STORAGE_KEY, JSON.stringify(preset));
60
+ return true;
61
+ } catch {
62
+ return false;
63
+ }
64
+ }
65
+
66
+ const CHANNEL_DEFAULT = {
67
+ prompt: '',
68
+ duration: 8,
69
+ durationMode: 'seconds',
70
+ bars: 4,
71
+ looping: true,
72
+ muted: false,
73
+ soloed: false,
74
+ knobs: { gain: -6, pan: 0, filter: 18000, delay: 0, reverb: 0 },
75
+ };
76
+
77
+ function defaultSession(channelCount) {
78
+ return {
79
+ bpm: 120,
80
+ launchQuantum: 4,
81
+ masterDb: 0,
82
+ injectBpm: true,
83
+ linkEnabled: false,
84
+ selectedModel: '',
85
+ selectedUnwrappedModel: '',
86
+ steps: 250,
87
+ randomSeed: true,
88
+ seedValue: '',
89
+ channels: Array.from({ length: channelCount }, () => ({
90
+ ...CHANNEL_DEFAULT,
91
+ knobs: { ...CHANNEL_DEFAULT.knobs },
92
+ })),
93
+ };
94
+ }
95
+
96
+ function loadSession(channelCount) {
97
+ const fallback = defaultSession(channelCount);
98
+ try {
99
+ const raw = localStorage.getItem(STORAGE_KEY);
100
+ if (!raw) return fallback;
101
+ const parsed = JSON.parse(raw);
102
+ // Merge against defaults so older saves don't crash on missing fields.
103
+ // Length shifts (channel count change between releases) are absorbed
104
+ // by always producing exactly `channelCount` channels.
105
+ const channels = Array.from({ length: channelCount }, (_, i) => ({
106
+ ...CHANNEL_DEFAULT,
107
+ ...(parsed.channels?.[i] || {}),
108
+ knobs: { ...CHANNEL_DEFAULT.knobs, ...(parsed.channels?.[i]?.knobs || {}) },
109
+ }));
110
+ return { ...fallback, ...parsed, channels };
111
+ } catch {
112
+ return fallback;
113
+ }
114
+ }
115
+
116
+ export function usePerformanceSession(channelCount = 4) {
117
+ const [session, setSession] = useState(() => loadSession(channelCount));
118
+ const persistTimerRef = useRef(null);
119
+
120
+ // Knobs and sliders fire many times per second; debounce writes so we don't
121
+ // hammer localStorage. Last-write-wins is fine for session continuity.
122
+ useEffect(() => {
123
+ if (persistTimerRef.current) clearTimeout(persistTimerRef.current);
124
+ persistTimerRef.current = setTimeout(() => {
125
+ try {
126
+ localStorage.setItem(STORAGE_KEY, JSON.stringify(session));
127
+ } catch { /* quota or serialization — non-fatal */ }
128
+ }, 250);
129
+ return () => {
130
+ if (persistTimerRef.current) clearTimeout(persistTimerRef.current);
131
+ };
132
+ }, [session]);
133
+
134
+ const updateGlobal = useCallback((key, value) => {
135
+ setSession(prev => (prev[key] === value ? prev : { ...prev, [key]: value }));
136
+ }, []);
137
+
138
+ const updateChannel = useCallback((index, partial) => {
139
+ setSession(prev => {
140
+ const channels = prev.channels.slice();
141
+ channels[index] = { ...channels[index], ...partial };
142
+ return { ...prev, channels };
143
+ });
144
+ }, []);
145
+
146
+ return { session, updateGlobal, updateChannel };
147
+ }
app/frontend/src/theme.js CHANGED
@@ -2119,6 +2119,28 @@ export const lossChartStyles = {
2119
  },
2120
  };
2121
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2122
  export const performancePanelStyles = {
2123
  root: {
2124
  display: 'flex',
@@ -2157,7 +2179,7 @@ export const performancePanelStyles = {
2157
  },
2158
  subtitle: {
2159
  color: 'text.secondary',
2160
- fontSize: '0.72rem',
2161
  },
2162
  headerPickers: {
2163
  display: 'flex',
@@ -2221,10 +2243,10 @@ export const performancePanelStyles = {
2221
  color,
2222
  }),
2223
  masterBadge: (color) => ({
2224
- fontSize: '0.72rem',
2225
  fontWeight: 600,
2226
  color,
2227
- letterSpacing: '0.14em',
2228
  px: 0.75,
2229
  py: 0.25,
2230
  borderRadius: 1,
@@ -2283,13 +2305,13 @@ export const performancePanelStyles = {
2283
  masterValue: {
2284
  textAlign: 'center',
2285
  color: 'primary.main',
2286
- fontSize: '0.68rem',
2287
  letterSpacing: '0.04em',
2288
  },
2289
  masterPeakValue: {
2290
  textAlign: 'center',
2291
  color: 'text.disabled',
2292
- fontSize: '0.58rem',
2293
  letterSpacing: '0.04em',
2294
  },
2295
  masterTransport: {
@@ -2303,7 +2325,7 @@ export const performancePanelStyles = {
2303
  masterBtn: (color, variant) => (theme) => ({
2304
  textTransform: 'none',
2305
  borderRadius: 1.5,
2306
- fontSize: '0.72rem',
2307
  py: 0.5,
2308
  ...(variant === 'play'
2309
  ? {
@@ -2348,10 +2370,10 @@ export const performanceChannelStyles = {
2348
  }),
2349
  channelBadge: (color) => ({
2350
  fontFamily: 'inherit',
2351
- fontSize: '0.78rem',
2352
  fontWeight: 600,
2353
  color,
2354
- letterSpacing: '0.08em',
2355
  px: 0.75,
2356
  py: 0.25,
2357
  borderRadius: 1,
@@ -2363,9 +2385,9 @@ export const performanceChannelStyles = {
2363
  gap: 0.5,
2364
  },
2365
  muteBtn: (active) => ({
2366
- width: 22,
2367
- height: 22,
2368
- fontSize: '0.65rem',
2369
  fontWeight: 700,
2370
  borderRadius: 1,
2371
  color: active ? '#fff' : 'text.secondary',
@@ -2377,9 +2399,9 @@ export const performanceChannelStyles = {
2377
  },
2378
  }),
2379
  soloBtn: (active) => ({
2380
- width: 22,
2381
- height: 22,
2382
- fontSize: '0.65rem',
2383
  fontWeight: 700,
2384
  borderRadius: 1,
2385
  color: active ? '#0c1018' : 'text.secondary',
@@ -2397,7 +2419,7 @@ export const performanceChannelStyles = {
2397
  },
2398
  promptField: (theme) => ({
2399
  '& .MuiOutlinedInput-root': {
2400
- fontSize: '0.75rem',
2401
  backgroundColor: theme.palette.mode === 'dark' ? 'rgba(9, 12, 18, 0.5)' : 'rgba(0, 0, 0, 0.04)',
2402
  borderRadius: 1.5,
2403
  '& textarea': { lineHeight: 1.3 },
@@ -2411,7 +2433,7 @@ export const performanceChannelStyles = {
2411
  },
2412
  durationLabel: {
2413
  fontFamily: 'inherit',
2414
- fontSize: '0.65rem',
2415
  color: 'text.secondary',
2416
  minWidth: 22,
2417
  },
@@ -2423,8 +2445,8 @@ export const performanceChannelStyles = {
2423
  }),
2424
  generateBtn: (color) => (theme) => ({
2425
  alignSelf: 'flex-end',
2426
- width: 28,
2427
- height: 28,
2428
  borderRadius: 1.5,
2429
  color,
2430
  border: `1px solid ${color}55`,
@@ -2449,8 +2471,8 @@ export const performanceChannelStyles = {
2449
  justifyContent: 'center',
2450
  color: 'text.disabled',
2451
  fontFamily: 'inherit',
2452
- fontSize: '0.65rem',
2453
- letterSpacing: '0.1em',
2454
  pointerEvents: 'none',
2455
  },
2456
  knobsGrid: {
@@ -2466,19 +2488,21 @@ export const performanceChannelStyles = {
2466
  gap: 0.25,
2467
  height: 70,
2468
  },
2469
- knobSlider: (color) => ({
2470
  height: 50,
2471
  color,
2472
- '& .MuiSlider-thumb': { width: 10, height: 10 },
2473
- '& .MuiSlider-rail': { opacity: 0.3, width: 2 },
2474
- '& .MuiSlider-track': { width: 2, border: 'none' },
 
 
2475
  }),
2476
  knobLabel: {
2477
  display: 'block',
2478
  fontFamily: 'inherit',
2479
- fontSize: '0.53rem',
2480
  color: 'text.secondary',
2481
- letterSpacing: '0.06em',
2482
  mt: 0.75,
2483
  },
2484
  transportRow: {
@@ -2491,8 +2515,8 @@ export const performanceChannelStyles = {
2491
  borderTopColor: 'divider',
2492
  },
2493
  transportBtn: (color, playing) => (theme) => ({
2494
- width: 26,
2495
- height: 26,
2496
  borderRadius: 1.5,
2497
  color: playing ? '#0c1018' : color,
2498
  backgroundColor: playing ? color : `${color}14`,
@@ -2501,8 +2525,8 @@ export const performanceChannelStyles = {
2501
  '&.Mui-disabled': theme.palette.mode === 'dark' ? { opacity: 0.3 } : {},
2502
  }),
2503
  loopBtn: (color, active) => ({
2504
- width: 22,
2505
- height: 22,
2506
  borderRadius: 1,
2507
  color: active ? color : 'text.secondary',
2508
  backgroundColor: active ? `${color}1F` : 'transparent',
 
2119
  },
2120
  };
2121
 
2122
+ // Shared visual tokens for the Performance page (panel, channels, MIDI menu).
2123
+ // One source of truth for the size/spacing/height scale so similar elements
2124
+ // match. The previous code carried 9+ distinct font sizes and 5+ letter
2125
+ // spacings across these surfaces; anything new should pick from this set.
2126
+ export const perfTokens = {
2127
+ fontSize: {
2128
+ knob: '0.58rem', // knob labels, pan label, master peak readout
2129
+ small: '0.66rem', // small labels: BPM unit, mute/solo, durationLabel, footer notes
2130
+ body: '0.72rem', // primary text: buttons, dropdowns, prompt field, mapping rows
2131
+ badge: '0.78rem', // section badges (MASTER, channel numbers)
2132
+ },
2133
+ letterSpacing: {
2134
+ wide: '0.08em', // uppercase labels and badges
2135
+ },
2136
+ height: {
2137
+ compact: 26, // primary compact controls (Link, MIDI, Q, BPM, transport, generate)
2138
+ sub: 22, // small subordinate square buttons (mute, solo, loop)
2139
+ },
2140
+ // Sharp 2px radius is the deliberate Ableton-style accent on Link/MIDI.
2141
+ // Everything else lives on the MUI scale via shape.borderRadius (= 1.5).
2142
+ };
2143
+
2144
  export const performancePanelStyles = {
2145
  root: {
2146
  display: 'flex',
 
2179
  },
2180
  subtitle: {
2181
  color: 'text.secondary',
2182
+ fontSize: perfTokens.fontSize.body,
2183
  },
2184
  headerPickers: {
2185
  display: 'flex',
 
2243
  color,
2244
  }),
2245
  masterBadge: (color) => ({
2246
+ fontSize: perfTokens.fontSize.badge,
2247
  fontWeight: 600,
2248
  color,
2249
+ letterSpacing: perfTokens.letterSpacing.wide,
2250
  px: 0.75,
2251
  py: 0.25,
2252
  borderRadius: 1,
 
2305
  masterValue: {
2306
  textAlign: 'center',
2307
  color: 'primary.main',
2308
+ fontSize: perfTokens.fontSize.small,
2309
  letterSpacing: '0.04em',
2310
  },
2311
  masterPeakValue: {
2312
  textAlign: 'center',
2313
  color: 'text.disabled',
2314
+ fontSize: perfTokens.fontSize.knob,
2315
  letterSpacing: '0.04em',
2316
  },
2317
  masterTransport: {
 
2325
  masterBtn: (color, variant) => (theme) => ({
2326
  textTransform: 'none',
2327
  borderRadius: 1.5,
2328
+ fontSize: perfTokens.fontSize.body,
2329
  py: 0.5,
2330
  ...(variant === 'play'
2331
  ? {
 
2370
  }),
2371
  channelBadge: (color) => ({
2372
  fontFamily: 'inherit',
2373
+ fontSize: perfTokens.fontSize.badge,
2374
  fontWeight: 600,
2375
  color,
2376
+ letterSpacing: perfTokens.letterSpacing.wide,
2377
  px: 0.75,
2378
  py: 0.25,
2379
  borderRadius: 1,
 
2385
  gap: 0.5,
2386
  },
2387
  muteBtn: (active) => ({
2388
+ width: perfTokens.height.sub,
2389
+ height: perfTokens.height.sub,
2390
+ fontSize: perfTokens.fontSize.small,
2391
  fontWeight: 700,
2392
  borderRadius: 1,
2393
  color: active ? '#fff' : 'text.secondary',
 
2399
  },
2400
  }),
2401
  soloBtn: (active) => ({
2402
+ width: perfTokens.height.sub,
2403
+ height: perfTokens.height.sub,
2404
+ fontSize: perfTokens.fontSize.small,
2405
  fontWeight: 700,
2406
  borderRadius: 1,
2407
  color: active ? '#0c1018' : 'text.secondary',
 
2419
  },
2420
  promptField: (theme) => ({
2421
  '& .MuiOutlinedInput-root': {
2422
+ fontSize: perfTokens.fontSize.body,
2423
  backgroundColor: theme.palette.mode === 'dark' ? 'rgba(9, 12, 18, 0.5)' : 'rgba(0, 0, 0, 0.04)',
2424
  borderRadius: 1.5,
2425
  '& textarea': { lineHeight: 1.3 },
 
2433
  },
2434
  durationLabel: {
2435
  fontFamily: 'inherit',
2436
+ fontSize: perfTokens.fontSize.small,
2437
  color: 'text.secondary',
2438
  minWidth: 22,
2439
  },
 
2445
  }),
2446
  generateBtn: (color) => (theme) => ({
2447
  alignSelf: 'flex-end',
2448
+ width: perfTokens.height.compact,
2449
+ height: perfTokens.height.compact,
2450
  borderRadius: 1.5,
2451
  color,
2452
  border: `1px solid ${color}55`,
 
2471
  justifyContent: 'center',
2472
  color: 'text.disabled',
2473
  fontFamily: 'inherit',
2474
+ fontSize: perfTokens.fontSize.small,
2475
+ letterSpacing: perfTokens.letterSpacing.wide,
2476
  pointerEvents: 'none',
2477
  },
2478
  knobsGrid: {
 
2488
  gap: 0.25,
2489
  height: 70,
2490
  },
2491
+ knobSlider: (color, fat = false) => ({
2492
  height: 50,
2493
  color,
2494
+ // Gain is visually fattened (wider track + bigger thumb) so it stands
2495
+ // out from LPF/DLY/REV — it's the dBFS-scaled "fader" of the four.
2496
+ '& .MuiSlider-thumb': { width: fat ? 12 : 10, height: fat ? 12 : 10 },
2497
+ '& .MuiSlider-rail': { opacity: 0.3, width: fat ? 4 : 2 },
2498
+ '& .MuiSlider-track': { width: fat ? 4 : 2, border: 'none' },
2499
  }),
2500
  knobLabel: {
2501
  display: 'block',
2502
  fontFamily: 'inherit',
2503
+ fontSize: perfTokens.fontSize.knob,
2504
  color: 'text.secondary',
2505
+ letterSpacing: perfTokens.letterSpacing.wide,
2506
  mt: 0.75,
2507
  },
2508
  transportRow: {
 
2515
  borderTopColor: 'divider',
2516
  },
2517
  transportBtn: (color, playing) => (theme) => ({
2518
+ width: perfTokens.height.compact,
2519
+ height: perfTokens.height.compact,
2520
  borderRadius: 1.5,
2521
  color: playing ? '#0c1018' : color,
2522
  backgroundColor: playing ? color : `${color}14`,
 
2525
  '&.Mui-disabled': theme.palette.mode === 'dark' ? { opacity: 0.3 } : {},
2526
  }),
2527
  loopBtn: (color, active) => ({
2528
+ width: perfTokens.height.sub,
2529
+ height: perfTokens.height.sub,
2530
  borderRadius: 1,
2531
  color: active ? color : 'text.secondary',
2532
  backgroundColor: active ? `${color}1F` : 'transparent',
app/frontend/src/utils/performanceAudio.js CHANGED
@@ -1,6 +1,3 @@
1
- // Default per-channel gain = -6 dBFS. Keeps the summed mix from clipping the
2
- // master limiter at four channels playing together, while leaving headroom to
3
- // push a single channel hot.
4
  export const DEFAULT_CHANNEL_GAIN = Math.pow(10, -6 / 20); // ≈ 0.5012
5
 
6
  let sharedCtx = null;
@@ -15,8 +12,6 @@ export function getAudioContext() {
15
  return sharedCtx;
16
  }
17
 
18
- // Real impulse responses served from /public/ir/. `id` is the stable key used
19
- // by the UI / setChannelImpulseResponse(); `file` is the actual filename.
20
  export const IMPULSE_RESPONSES = [
21
  { id: 'hall', name: 'Opera Hall', file: 'Scala Milan Opera Hall.wav' },
22
  { id: 'room', name: 'Drum Room', file: 'Nice Drum Room.wav' },
@@ -53,9 +48,6 @@ export function getImpulseResponseBuffer(id) {
53
  return irBufferCache.get(id);
54
  }
55
 
56
- // Synthetic hall IR: early reflections + diffuse noise tail with progressive
57
- // high-frequency damping and per-channel decorrelation. Swap this out by
58
- // dropping a real IR WAV into /public/ir/ and fetching it into a ConvolverNode.
59
  const EARLY_REFLECTIONS_MS = [
60
  [7, 0.55], [13, -0.42], [19, 0.36], [28, -0.30],
61
  [41, 0.26], [56, 0.22], [73, -0.18], [91, 0.15],
@@ -79,15 +71,12 @@ function getImpulse(ctx, duration = 2.8, decaySeconds = 1.6, damping = 0.55) {
79
  if (idx < length) data[idx] += amp * 0.8;
80
  }
81
 
82
- // Diffuse tail: white noise gated by exp decay, through a 1-pole LP
83
- // whose cutoff shrinks over time so high freqs die first (natural air).
84
  let lpState = 0;
85
  const predelaySec = 0.012;
86
  for (let i = 0; i < length; i++) {
87
  const t = i / sr;
88
  if (t < predelaySec) continue;
89
  const env = Math.exp(-(t - predelaySec) / decaySeconds);
90
- // Damping grows with time: alpha goes from (1 - damping*0.4) → (1 - damping*0.95)
91
  const progression = Math.min(1, (t - predelaySec) / decaySeconds);
92
  const alpha = 1 - damping * (0.4 + 0.55 * progression);
93
  const noise = (Math.random() * 2 - 1) * env;
@@ -127,9 +116,6 @@ export class ChannelStrip {
127
 
128
  this.dryGain = ctx.createGain();
129
  this.dryGain.gain.value = 1.0;
130
-
131
- // Delay time is tempo-locked to an 8th note. Default seed = 120 BPM
132
- // (0.25 s); PerformanceEngine overwrites this once the panel sets BPM.
133
  this.delayNode = ctx.createDelay(2.0);
134
  this.delayNode.delayTime.value = 0.25;
135
  this.delayFeedback = ctx.createGain();
@@ -146,8 +132,6 @@ export class ChannelStrip {
146
  this.channelGain.gain.value = DEFAULT_CHANNEL_GAIN;
147
  this._lastUserGain = DEFAULT_CHANNEL_GAIN;
148
 
149
- // Glue compressor on the post-mix channel bus. Gentle defaults —
150
- // evens out transients without obvious pumping.
151
  this.compressor = ctx.createDynamicsCompressor();
152
  this.compressor.threshold.value = -16;
153
  this.compressor.knee.value = 8;
@@ -187,7 +171,7 @@ export class ChannelStrip {
187
  this.buffer = await this.ctx.decodeAudioData(arrayBuffer);
188
  }
189
 
190
- play(loop = this.isLooping) {
191
  if (!this.buffer) return;
192
  this.stop();
193
  this.isLooping = loop;
@@ -201,7 +185,8 @@ export class ChannelStrip {
201
  this.isPlaying = false;
202
  }
203
  };
204
- src.start(0);
 
205
  this.source = src;
206
  this.isPlaying = true;
207
  }
@@ -224,8 +209,6 @@ export class ChannelStrip {
224
  if (buffer) this.reverbNode.buffer = buffer;
225
  }
226
  setDelayTimeForBpm(bpm) {
227
- // 8th note in seconds = 60 / bpm / 2 = 30 / bpm.
228
- // Glide over ~40 ms so BPM edits slide rather than click.
229
  const safeBpm = Math.max(1, bpm);
230
  const eighthSec = Math.min(30 / safeBpm, 2.0);
231
  this.delayNode.delayTime.setTargetAtTime(
@@ -310,11 +293,6 @@ export class PerformanceEngine {
310
  this.ctx = ctx;
311
  this.masterBus = ctx.createGain();
312
  this.masterBus.gain.value = 0.9;
313
-
314
- // Master limiter: brick-wall style defaults to catch sums that push
315
- // past 0 dBFS when many channels play at once. DynamicsCompressor
316
- // with a high ratio + fast attack is the best built-in approximation
317
- // of a look-ahead limiter that Web Audio offers.
318
  this.masterLimiter = ctx.createDynamicsCompressor();
319
  this.masterLimiter.threshold.value = -1.0;
320
  this.masterLimiter.knee.value = 0;
@@ -330,9 +308,20 @@ export class PerformanceEngine {
330
  this.masterAnalyser.connect(ctx.destination);
331
  this.channels = Array.from({ length: channelCount }, () => new ChannelStrip(this.masterBus));
332
 
333
- // Load real IRs asynchronously and swap them in when ready. Until
334
- // they arrive (or if loading fails), channels keep using the
335
- // synthetic fallback set in ChannelStrip's constructor.
 
 
 
 
 
 
 
 
 
 
 
336
  this.currentImpulseId = DEFAULT_IR_ID;
337
  loadImpulseResponses(ctx).then(() => {
338
  const buf = getImpulseResponseBuffer(this.currentImpulseId);
@@ -356,9 +345,79 @@ export class PerformanceEngine {
356
  }
357
 
358
  setBpm(bpm) {
 
 
 
 
 
 
 
 
 
 
 
 
 
359
  this.channels.forEach(ch => ch.setDelayTimeForBpm(bpm));
360
  }
361
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
362
  setMasterGain(value) {
363
  this.masterBus.gain.setTargetAtTime(value, this.ctx.currentTime, 0.01);
364
  }
@@ -389,7 +448,8 @@ export class PerformanceEngine {
389
  }
390
 
391
  playAll(loop = true) {
392
- this.channels.forEach(ch => { if (ch.buffer) ch.play(loop); });
 
393
  }
394
 
395
  stopAll() {
 
 
 
 
1
  export const DEFAULT_CHANNEL_GAIN = Math.pow(10, -6 / 20); // ≈ 0.5012
2
 
3
  let sharedCtx = null;
 
12
  return sharedCtx;
13
  }
14
 
 
 
15
  export const IMPULSE_RESPONSES = [
16
  { id: 'hall', name: 'Opera Hall', file: 'Scala Milan Opera Hall.wav' },
17
  { id: 'room', name: 'Drum Room', file: 'Nice Drum Room.wav' },
 
48
  return irBufferCache.get(id);
49
  }
50
 
 
 
 
51
  const EARLY_REFLECTIONS_MS = [
52
  [7, 0.55], [13, -0.42], [19, 0.36], [28, -0.30],
53
  [41, 0.26], [56, 0.22], [73, -0.18], [91, 0.15],
 
71
  if (idx < length) data[idx] += amp * 0.8;
72
  }
73
 
 
 
74
  let lpState = 0;
75
  const predelaySec = 0.012;
76
  for (let i = 0; i < length; i++) {
77
  const t = i / sr;
78
  if (t < predelaySec) continue;
79
  const env = Math.exp(-(t - predelaySec) / decaySeconds);
 
80
  const progression = Math.min(1, (t - predelaySec) / decaySeconds);
81
  const alpha = 1 - damping * (0.4 + 0.55 * progression);
82
  const noise = (Math.random() * 2 - 1) * env;
 
116
 
117
  this.dryGain = ctx.createGain();
118
  this.dryGain.gain.value = 1.0;
 
 
 
119
  this.delayNode = ctx.createDelay(2.0);
120
  this.delayNode.delayTime.value = 0.25;
121
  this.delayFeedback = ctx.createGain();
 
132
  this.channelGain.gain.value = DEFAULT_CHANNEL_GAIN;
133
  this._lastUserGain = DEFAULT_CHANNEL_GAIN;
134
 
 
 
135
  this.compressor = ctx.createDynamicsCompressor();
136
  this.compressor.threshold.value = -16;
137
  this.compressor.knee.value = 8;
 
171
  this.buffer = await this.ctx.decodeAudioData(arrayBuffer);
172
  }
173
 
174
+ play(loop = this.isLooping, startTime = 0) {
175
  if (!this.buffer) return;
176
  this.stop();
177
  this.isLooping = loop;
 
185
  this.isPlaying = false;
186
  }
187
  };
188
+
189
+ src.start(Math.max(0, startTime));
190
  this.source = src;
191
  this.isPlaying = true;
192
  }
 
209
  if (buffer) this.reverbNode.buffer = buffer;
210
  }
211
  setDelayTimeForBpm(bpm) {
 
 
212
  const safeBpm = Math.max(1, bpm);
213
  const eighthSec = Math.min(30 / safeBpm, 2.0);
214
  this.delayNode.delayTime.setTargetAtTime(
 
293
  this.ctx = ctx;
294
  this.masterBus = ctx.createGain();
295
  this.masterBus.gain.value = 0.9;
 
 
 
 
 
296
  this.masterLimiter = ctx.createDynamicsCompressor();
297
  this.masterLimiter.threshold.value = -1.0;
298
  this.masterLimiter.knee.value = 0;
 
308
  this.masterAnalyser.connect(ctx.destination);
309
  this.channels = Array.from({ length: channelCount }, () => new ChannelStrip(this.masterBus));
310
 
311
+
312
+ this.linkSnapshot = null;
313
+ this.launchQuantum = 0;
314
+
315
+ // Internal transport: an always-running beat clock anchored in audio
316
+ // time. Used for launch quantization when Ableton Link isn't active,
317
+ // so 'Q' still lines launches up to the bar even with no peer.
318
+ // BPM changes rebase the anchor (see setBpm) so phase is preserved.
319
+ this.internalTransport = {
320
+ originAudioTime: ctx.currentTime,
321
+ anchorBeat: 0,
322
+ bpm: 120,
323
+ };
324
+
325
  this.currentImpulseId = DEFAULT_IR_ID;
326
  loadImpulseResponses(ctx).then(() => {
327
  const buf = getImpulseResponseBuffer(this.currentImpulseId);
 
345
  }
346
 
347
  setBpm(bpm) {
348
+ const safe = Number(bpm);
349
+ if (Number.isFinite(safe) && safe > 0) {
350
+ // Rebase the internal transport so its current beat position is
351
+ // preserved across the BPM change. Without rebasing, switching
352
+ // 120→140 would jump the next-quantized beat by minutes' worth
353
+ // of time in the wrong direction.
354
+ const tr = this.internalTransport;
355
+ const now = this.ctx.currentTime;
356
+ const elapsed = now - tr.originAudioTime;
357
+ tr.anchorBeat += elapsed * tr.bpm / 60;
358
+ tr.originAudioTime = now;
359
+ tr.bpm = safe;
360
+ }
361
  this.channels.forEach(ch => ch.setDelayTimeForBpm(bpm));
362
  }
363
 
364
+ setLinkSnapshot(snapshot) {
365
+ this.linkSnapshot = snapshot;
366
+ }
367
+
368
+ setLaunchQuantum(beats) {
369
+ const v = Number(beats);
370
+ this.launchQuantum = Number.isFinite(v) && v > 0 ? v : 0;
371
+ }
372
+
373
+ getNextQuantizedAudioTime() {
374
+ const quantum = this.launchQuantum;
375
+ if (!quantum) return 0;
376
+
377
+ // First-launch shortcut: if nothing is currently playing, fire
378
+ // immediately and (re)anchor the internal transport at "now". This
379
+ // matches Live's Session View — the user pressed Play, they expect
380
+ // audio, not silence until the next bar. Subsequent launches see at
381
+ // least one channel playing and quantize as normal. Link's clock is
382
+ // external so we leave its snapshot alone.
383
+ const anythingPlaying = this.channels.some(c => c.isPlaying);
384
+ if (!anythingPlaying) {
385
+ if (!this.linkSnapshot) {
386
+ this.internalTransport.originAudioTime = this.ctx.currentTime;
387
+ this.internalTransport.anchorBeat = 0;
388
+ }
389
+ return 0;
390
+ }
391
+
392
+ // Prefer the Link snapshot when active so the app stays in phase with
393
+ // external peers. Falls through to the internal transport when Link
394
+ // isn't running so 'Q' still works standalone.
395
+ const snap = this.linkSnapshot;
396
+ if (snap && snap.bpm) {
397
+ const elapsedSec = (performance.now() - snap.capturedAt) / 1000;
398
+ const currentBeat = snap.beat + elapsedSec * (snap.bpm / 60);
399
+ let nextBeat = Math.ceil(currentBeat / quantum) * quantum;
400
+ if (nextBeat - currentBeat < 1e-6) nextBeat += quantum;
401
+ const secondsUntil = (nextBeat - currentBeat) * 60 / snap.bpm;
402
+ return this.ctx.currentTime + secondsUntil;
403
+ }
404
+
405
+ const tr = this.internalTransport;
406
+ if (!tr.bpm) return 0;
407
+ const elapsedSec = this.ctx.currentTime - tr.originAudioTime;
408
+ const currentBeat = tr.anchorBeat + elapsedSec * (tr.bpm / 60);
409
+ let nextBeat = Math.ceil(currentBeat / quantum) * quantum;
410
+ if (nextBeat - currentBeat < 1e-6) nextBeat += quantum;
411
+ const secondsUntil = (nextBeat - currentBeat) * 60 / tr.bpm;
412
+ return this.ctx.currentTime + secondsUntil;
413
+ }
414
+
415
+ playChannel(index, loop) {
416
+ const ch = this.channels[index];
417
+ if (!ch || !ch.buffer) return;
418
+ ch.play(loop, this.getNextQuantizedAudioTime());
419
+ }
420
+
421
  setMasterGain(value) {
422
  this.masterBus.gain.setTargetAtTime(value, this.ctx.currentTime, 0.01);
423
  }
 
448
  }
449
 
450
  playAll(loop = true) {
451
+ const startTime = this.getNextQuantizedAudioTime();
452
+ this.channels.forEach(ch => { if (ch.buffer) ch.play(loop, startTime); });
453
  }
454
 
455
  stopAll() {
models/config/model_config_small.json CHANGED
@@ -87,6 +87,7 @@
87
  "cond_dim": 768
88
  },
89
  "diffusion": {
 
90
  "cross_attention_cond_ids": [
91
  "prompt",
92
  "seconds_total"
 
87
  "cond_dim": 768
88
  },
89
  "diffusion": {
90
+ "diffusion_objective": "rectified_flow",
91
  "cross_attention_cond_ids": [
92
  "prompt",
93
  "seconds_total"
stable-audio-tools/stable_audio_tools/inference/sampling.py CHANGED
@@ -134,6 +134,30 @@ def sample_rk4(model, x, steps, sigma_max=1, callback=None, dist_shift=None, **e
134
  # If we are on the last timestep, output the denoised data
135
  return x
136
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
137
  @torch.no_grad()
138
  def sample_flow_dpmpp(model, x, steps, sigma_max=1, callback=None, dist_shift=None, **extra_args):
139
  """Draws samples from a model given starting noise. DPM-Solver++ for RF models"""
@@ -371,4 +395,6 @@ def sample_rf(
371
  elif sampler_type == "rk4":
372
  return sample_rk4(model_fn, x, steps, sigma_max, callback=callback, **extra_args)
373
  elif sampler_type == "dpmpp":
374
- return sample_flow_dpmpp(model_fn, x, steps, sigma_max, callback=callback, **extra_args)
 
 
 
134
  # If we are on the last timestep, output the denoised data
135
  return x
136
 
137
+ @torch.no_grad()
138
+ def sample_flow_pingpong(model, x, steps, sigma_max=1, callback=None, dist_shift=None, **extra_args):
139
+ """Draws samples from a model given starting noise. Ping-pong sampling for distilled RF models."""
140
+
141
+ # Make tensor of ones to broadcast the single t values
142
+ ts = x.new_ones([x.shape[0]])
143
+
144
+ # Create the noise schedule
145
+ t = torch.linspace(sigma_max, 0, steps + 1)
146
+
147
+ if dist_shift is not None:
148
+ t = dist_shift.time_shift(t, x.shape[-1])
149
+
150
+ for i in trange(len(t) - 1, disable=False):
151
+ denoised = x - t[i] * model(x, t[i] * ts, **extra_args)
152
+ if callback is not None:
153
+ callback({'x': x, 'i': i, 't': t[i], 'sigma': t[i], 'sigma_hat': t[i], 'denoised': denoised})
154
+
155
+ t_next = t[i + 1]
156
+ x = (1 - t_next) * denoised + t_next * torch.randn_like(x)
157
+
158
+ return x
159
+
160
+
161
  @torch.no_grad()
162
  def sample_flow_dpmpp(model, x, steps, sigma_max=1, callback=None, dist_shift=None, **extra_args):
163
  """Draws samples from a model given starting noise. DPM-Solver++ for RF models"""
 
395
  elif sampler_type == "rk4":
396
  return sample_rk4(model_fn, x, steps, sigma_max, callback=callback, **extra_args)
397
  elif sampler_type == "dpmpp":
398
+ return sample_flow_dpmpp(model_fn, x, steps, sigma_max, callback=callback, **extra_args)
399
+ elif sampler_type == "pingpong":
400
+ return sample_flow_pingpong(model_fn, x, steps, sigma_max, callback=callback, **extra_args)