MazCodes commited on
Commit
0573fbf
·
verified ·
1 Parent(s): 8206e1c

Upload folder using huggingface_hub

Browse files
Files changed (47) hide show
  1. .dockerignore +113 -25
  2. .gitattributes +2 -0
  3. Dockerfile +9 -19
  4. README.md +1 -1
  5. app/backend/app.py +418 -245
  6. app/backend/data/auto_annotator.py +406 -0
  7. app/core/config.py +6 -1
  8. app/core/generation/audio_generator.py +42 -40
  9. app/core/model_manager.py +16 -44
  10. app/core/training/lora_trainer.py +68 -47
  11. app/frontend/build/BitcountSingle-VariableFont_CRSV,ELSH,ELXP,slnt,wght.ttf +3 -0
  12. app/frontend/build/assets/index-D-qgc0vE.js +0 -0
  13. app/frontend/build/favicon.ico +0 -0
  14. app/frontend/build/fragmenta_background.png +3 -0
  15. app/frontend/build/fragmenta_icon_1024.png +0 -0
  16. app/frontend/build/index.html +35 -0
  17. app/frontend/build/manifest.json +15 -0
  18. app/frontend/index.html +35 -0
  19. app/frontend/package-lock.json +0 -0
  20. app/frontend/package.json +11 -26
  21. app/frontend/public/favicon.ico +2 -2
  22. app/frontend/public/fragmenta_icon_1024.png +2 -2
  23. app/frontend/src/App.js +0 -0
  24. app/frontend/src/api.js +41 -0
  25. app/frontend/src/components/AudioUploadRow.js +104 -0
  26. app/frontend/src/components/BulkAnnotatePanel.js +316 -0
  27. app/frontend/src/components/CheckpointManager.js +211 -0
  28. app/frontend/src/components/GeneratedFragmentsWindow.js +114 -0
  29. app/frontend/src/components/HfAuthDialog.js +245 -0
  30. app/frontend/src/components/LossChart.js +104 -0
  31. app/frontend/src/components/ModelUnwrapButton.js +58 -0
  32. app/frontend/src/components/TabPanel.js +21 -0
  33. app/frontend/src/components/TrainingMonitor.js +123 -0
  34. app/frontend/src/components/WelcomePage.js +102 -0
  35. app/frontend/src/index.js +0 -8
  36. app/frontend/src/theme.js +2113 -0
  37. app/frontend/src/utils/format.js +8 -0
  38. app/frontend/vite.config.js +26 -0
  39. docker-entrypoint.sh +12 -3
  40. models/config/dataset-config.json +4 -3
  41. requirements.txt +11 -23
  42. stable-audio-tools/custom_metadata.py +12 -12
  43. stable-audio-tools/package-lock.json +6 -0
  44. stable-audio-tools/stable_audio_tools/models/conditioners.py +6 -3
  45. stable-audio-tools/stable_audio_tools/training/arc.py +1 -1
  46. stable-audio-tools/stable_audio_tools/training/diffusion.py +18 -18
  47. stable-audio-tools/train.py +1 -1
.dockerignore CHANGED
@@ -1,44 +1,132 @@
1
- # Version control
2
  .git
3
  .gitignore
4
 
5
- # Python
6
- venv/
7
  __pycache__/
8
- *.pyc
9
- *.pyo
10
- *.egg-info
11
- .pytest_cache/
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
12
 
13
- # Node
14
- app/frontend/node_modules/
 
 
15
 
16
- # Build artifacts
17
- build/
 
 
 
 
 
18
  dist/
19
  distribution/
 
 
 
 
 
 
 
 
20
  *.spec
 
 
 
 
 
 
 
 
 
21
 
22
- # IDE
 
 
 
 
 
 
 
23
  .vscode/
24
  .idea/
25
  *.swp
26
  *.swo
27
-
28
- # OS files
29
  .DS_Store
 
 
 
 
 
30
  Thumbs.db
31
 
32
- # Large model files (mounted as volumes instead)
33
- models/pretrained/*.safetensors
34
- models/fine_tuned/
35
-
36
- # User state — mounted as volumes, must NOT be baked into the image
37
- config/
38
- data/
39
- logs/
 
40
 
41
- # Docker
42
- Dockerfile*
43
- docker-compose*.yml
 
44
  DOCKER.md
 
 
 
 
 
1
+ # ── Version control ───────────────────────────────────────────────────────────
2
  .git
3
  .gitignore
4
 
5
+ # ── Python ────────────────────────────────────────────────────────────────────
 
6
  __pycache__/
7
+ *.py[cod]
8
+ *$py.class
9
+ *.so
10
+ .Python
11
+ develop-eggs/
12
+ downloads/
13
+ eggs/
14
+ .eggs/
15
+ lib/
16
+ lib64/
17
+ parts/
18
+ sdist/
19
+ var/
20
+ wheels/
21
+ share/python-wheels/
22
+ *.egg-info/
23
+ .installed.cfg
24
+ *.egg
25
+ MANIFEST
26
+ venv/
27
+ env/
28
+ ENV/
29
+ env.bak/
30
+ venv.bak/
31
+
32
+ # ── Node / frontend build tooling ────────────────────────────────────────────
33
+ node_modules/
34
+ npm-debug.log*
35
+ yarn-debug.log*
36
+ yarn-error.log*
37
+ .npm
38
+ .eslintcache
39
+
40
+ # ── Models (downloaded at runtime, never baked into the image) ───────────────
41
+ models/pretrained/*.safetensors
42
+ models/pretrained/*.ckpt
43
+ models/pretrained/*.pth
44
+ models/pretrained/clap
45
+
46
+ # ── Audio / output (runtime data) ────────────────────────────────────────────
47
+ output/
48
+ *.wav
49
+ *.mp3
50
+ *.flac
51
+ *.aiff
52
+ *.ogg
53
+
54
+ # ── Logs ─────────────────────────────────────────────────────────────────────
55
+ logs/
56
+ *.log
57
+ app/backend/*.log
58
+ lightning_logs/
59
 
60
+ # ── Checkpoints ──────────────────────────────────────────────────────────────
61
+ *.ckpt
62
+ models/fine_tuned/*/
63
+ !models/fine_tuned/.gitkeep
64
 
65
+ # ── Config (user-generated at runtime) ───────────────────────────────────────
66
+ config/terms_accepted.json
67
+
68
+ # ── Distribution / build artefacts ───────────────────────────────────────────
69
+ build/fragmenta/
70
+ build/fragmenta.spec
71
+ build/windows_source/
72
  dist/
73
  distribution/
74
+ *.dmg
75
+ *.app
76
+ *.exe
77
+ *.iss
78
+ *.pkg
79
+ *.deb
80
+ *.rpm
81
+ *.AppImage
82
  *.spec
83
+ scripts/__pycache__/
84
+
85
+ # ── Temporary files ───────────────────────────────────────────────────────────
86
+ tmp/
87
+ temp/
88
+ *.temp
89
+ *.tmp
90
+ .cache
91
+ .parcel-cache
92
 
93
+ # ── Environment / secrets ─────────────────────────────────────────────────────
94
+ .env
95
+ .env.local
96
+ .env.development.local
97
+ .env.test.local
98
+ .env.production.local
99
+
100
+ # ── IDE / OS ──────────────────────────────────────────────────────────────────
101
  .vscode/
102
  .idea/
103
  *.swp
104
  *.swo
105
+ *~
 
106
  .DS_Store
107
+ .DS_Store?
108
+ ._*
109
+ .Spotlight-V100
110
+ .Trashes
111
+ ehthumbs.db
112
  Thumbs.db
113
 
114
+ # ── Desktop launcher scripts (not needed in Docker) ───────────────────────────
115
+ fragmenta.sh
116
+ fragmenta.bat
117
+ fragmenta.command
118
+ scripts/
119
+ start.py
120
+ main.py
121
+ debug_binaries.py
122
+ SIMPLIFICATION_PLAN.md
123
 
124
+ # ── Docs (not needed in image) ────────────────────────────────────────────────
125
+ README.md
126
+ NOTICE.md
127
+ DOCKERHUB.md
128
  DOCKER.md
129
+
130
+ # ── Data (mounted as volume at runtime) ──────────────────────────────────────
131
+ /data/
132
+ .claude
.gitattributes CHANGED
@@ -37,3 +37,5 @@ app/frontend/public/BitcountSingle-VariableFont_CRSV,ELSH,ELXP,slnt,wght.ttf fil
37
  app/frontend/public/favicon.ico filter=lfs diff=lfs merge=lfs -text
38
  app/frontend/public/fragmenta_background.png filter=lfs diff=lfs merge=lfs -text
39
  app/frontend/public/fragmenta_icon_1024.png filter=lfs diff=lfs merge=lfs -text
 
 
 
37
  app/frontend/public/favicon.ico filter=lfs diff=lfs merge=lfs -text
38
  app/frontend/public/fragmenta_background.png filter=lfs diff=lfs merge=lfs -text
39
  app/frontend/public/fragmenta_icon_1024.png filter=lfs diff=lfs merge=lfs -text
40
+ app/frontend/build/BitcountSingle-VariableFont_CRSV,ELSH,ELXP,slnt,wght.ttf filter=lfs diff=lfs merge=lfs -text
41
+ app/frontend/build/fragmenta_background.png filter=lfs diff=lfs merge=lfs -text
Dockerfile CHANGED
@@ -7,23 +7,12 @@
7
  # CPU Spaces: torch.cuda.is_available() → False → runs on CPU
8
  # GPU Spaces: NVIDIA runtime injected by HF → torch.cuda.is_available() → True
9
  #
 
 
 
10
  # Port: 7860 (HF Spaces requirement)
11
  # =============================================================================
12
 
13
- # ---------------------------------------------------------------------------
14
- # Stage 1: Build the React frontend
15
- # ---------------------------------------------------------------------------
16
- FROM node:20-slim AS frontend-builder
17
-
18
- WORKDIR /build/frontend
19
- COPY app/frontend/package.json app/frontend/package-lock.json* ./
20
- RUN npm ci --prefer-offline --no-audit 2>/dev/null || npm install
21
- COPY app/frontend/ ./
22
- RUN npm run build
23
-
24
- # ---------------------------------------------------------------------------
25
- # Stage 2: Python backend (CPU + GPU via bundled CUDA in PyTorch wheels)
26
- # ---------------------------------------------------------------------------
27
  FROM python:3.11-slim-bookworm
28
 
29
  ENV DEBIAN_FRONTEND=noninteractive
@@ -36,7 +25,7 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
36
  curl \
37
  && apt-get clean && rm -rf /var/lib/apt/lists/*
38
 
39
- RUN pip install --no-cache-dir --root-user-action=ignore --upgrade pip setuptools wheel
40
 
41
  WORKDIR /app
42
 
@@ -50,16 +39,16 @@ COPY requirements.txt .
50
  # On CPU Spaces it gracefully falls back to CPU — no extra cost.
51
  RUN pip install --no-cache-dir --root-user-action=ignore torch torchvision torchaudio
52
 
53
- # Install remaining requirements (skip desktop-only & CUDA index lines)
54
- RUN grep -ivE 'pyqt6|pyqt6-webengine|flash-attn|extra-index-url' requirements.txt > requirements_docker.txt \
 
55
  && pip install --no-cache-dir --root-user-action=ignore -r requirements_docker.txt \
56
  && rm requirements_docker.txt
57
 
58
  # ---------------------------------------------------------------------------
59
- # Application code
60
  # ---------------------------------------------------------------------------
61
  COPY . .
62
- COPY --from=frontend-builder /build/frontend/build ./app/frontend/build
63
 
64
  # Install stable-audio-tools in-tree
65
  RUN pip install --no-cache-dir --root-user-action=ignore -e ./stable-audio-tools/
@@ -109,6 +98,7 @@ ENV FRAGMENTA_LOG_LEVEL=INFO
109
  ENV FRAGMENTA_DOCKER=1
110
  ENV HOME=/home/user
111
  ENV PATH="/home/user/.local/bin:${PATH}"
 
112
 
113
  EXPOSE 7860
114
 
 
7
  # CPU Spaces: torch.cuda.is_available() → False → runs on CPU
8
  # GPU Spaces: NVIDIA runtime injected by HF → torch.cuda.is_available() → True
9
  #
10
+ # The React frontend is pre-built and committed under app/frontend/build/,
11
+ # so there is no Node.js build step here.
12
+ #
13
  # Port: 7860 (HF Spaces requirement)
14
  # =============================================================================
15
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
16
  FROM python:3.11-slim-bookworm
17
 
18
  ENV DEBIAN_FRONTEND=noninteractive
 
25
  curl \
26
  && apt-get clean && rm -rf /var/lib/apt/lists/*
27
 
28
+ RUN pip install --no-cache-dir --root-user-action=ignore --upgrade pip wheel
29
 
30
  WORKDIR /app
31
 
 
39
  # On CPU Spaces it gracefully falls back to CPU — no extra cost.
40
  RUN pip install --no-cache-dir --root-user-action=ignore torch torchvision torchaudio
41
 
42
+ # Install remaining requirements. Skip desktop-only deps (PyQt6, pywebview's
43
+ # GTK bindings Pycairo/PyGObject), flash-attn, and CUDA index URL.
44
+ RUN grep -ivE 'pyqt6|pyqt6-webengine|flash-attn|pycairo|pygobject|pywebview|extra-index-url' requirements.txt > requirements_docker.txt \
45
  && pip install --no-cache-dir --root-user-action=ignore -r requirements_docker.txt \
46
  && rm requirements_docker.txt
47
 
48
  # ---------------------------------------------------------------------------
49
+ # Application code (includes pre-built frontend under app/frontend/build/)
50
  # ---------------------------------------------------------------------------
51
  COPY . .
 
52
 
53
  # Install stable-audio-tools in-tree
54
  RUN pip install --no-cache-dir --root-user-action=ignore -e ./stable-audio-tools/
 
98
  ENV FRAGMENTA_DOCKER=1
99
  ENV HOME=/home/user
100
  ENV PATH="/home/user/.local/bin:${PATH}"
101
+ ENV OMP_NUM_THREADS=4
102
 
103
  EXPOSE 7860
104
 
README.md CHANGED
@@ -13,7 +13,7 @@ license: apache-2.0
13
 
14
  Generate and fine-tune audio from text prompts using Stable Audio Open.
15
 
16
- **Hardware:** cpu-upgrade (CPU)
17
 
18
  ## Getting Started
19
 
 
13
 
14
  Generate and fine-tune audio from text prompts using Stable Audio Open.
15
 
16
+ **Hardware:** cpu-basic (CPU)
17
 
18
  ## Getting Started
19
 
app/backend/app.py CHANGED
@@ -2,6 +2,9 @@ from utils.validators import Validator
2
  from utils.exceptions import ModelNotFoundError, ValidationError, GenerationError
3
  from utils.api_responses import APIResponse, handle_api_error
4
  from utils.logger import setup_logging, get_logger
 
 
 
5
  from app.core.config import get_config
6
  from flask import Flask, request, jsonify, send_file, send_from_directory
7
  from flask_cors import CORS
@@ -14,13 +17,6 @@ import json
14
  import logging
15
  from werkzeug.serving import WSGIRequestHandler
16
 
17
- # Heavy imports deferred to _ensure_components() for fast startup
18
- AudioGenerator = None
19
- SimpleAudioProcessor = None
20
- start_training_func = None
21
- get_training_status = None
22
- stop_training = None
23
-
24
  sys.path.append(os.path.abspath(
25
  os.path.join(os.path.dirname(__file__), '../../')))
26
 
@@ -62,6 +58,13 @@ def request_entity_too_large(error):
62
 
63
  DEBUG_MODE = os.environ.get('FRAGMENTA_DEBUG', 'false').lower() == 'true'
64
 
 
 
 
 
 
 
 
65
  config = None
66
  audio_processor = None
67
  generator = None
@@ -71,40 +74,17 @@ _init_error = None
71
 
72
 
73
  def _ensure_components():
74
- """Initialise backend components on first use. Thread-safe.
75
-
76
- Allows retry: if initialisation failed previously, it will attempt
77
- again (e.g. after the user downloads a model or fixes a dependency).
78
- """
79
  global config, audio_processor, generator, model_manager
80
  global _components_initialised, _init_error
81
 
82
  if _components_initialised:
83
  return
84
-
 
85
 
86
  try:
87
  logger.info("Initializing Backend API components (lazy)…")
88
-
89
- # Deferred heavy imports — keeps startup fast
90
- global AudioGenerator, SimpleAudioProcessor
91
- global start_training_func, get_training_status, stop_training
92
- if AudioGenerator is None:
93
- from app.core.generation.audio_generator import AudioGenerator as _AG
94
- AudioGenerator = _AG
95
- if SimpleAudioProcessor is None:
96
- from app.backend.data.simple_audio_processor import SimpleAudioProcessor as _SAP
97
- SimpleAudioProcessor = _SAP
98
- if start_training_func is None:
99
- from app.core.training.lora_trainer import (
100
- start_training as _st,
101
- get_training_status as _gts,
102
- stop_training as _stop,
103
- )
104
- start_training_func = _st
105
- get_training_status = _gts
106
- stop_training = _stop
107
-
108
  config = get_config()
109
  audio_processor = SimpleAudioProcessor(
110
  model_config_path=config.get_path("models_config") / "model_config.json"
@@ -115,7 +95,6 @@ def _ensure_components():
115
  model_manager = ModelManager(config)
116
 
117
  _components_initialised = True
118
- _init_error = None
119
  logger.info("Backend components initialized successfully")
120
 
121
  except Exception as e:
@@ -123,45 +102,18 @@ def _ensure_components():
123
  logger.error(f"Failed to initialize backend components: {e}")
124
  raise
125
 
126
- _download_progress = {}
127
-
128
- _LAZY_INIT_EXEMPT_PATHS = {
129
- '/api/health',
130
- '/api/environment',
131
- '/api/welcome-page-closed',
132
- '/api/welcome-page-status',
133
- '/api/open-output-folder',
134
- '/api/open-documentation',
135
- '/api/base-models/status',
136
- '/api/output-files',
137
- '/api/hf-token',
138
- '/api/hf-token/status',
139
- '/api/process-files',
140
- '/api/status',
141
- '/api/models',
142
- '/api/models-status',
143
- '/api/gpu-memory-status',
144
- '/api/start-fresh',
145
- '/api/license-info',
146
- '/api/debug-status',
147
- '/api/toggle-debug',
148
- }
149
-
150
 
151
  @app.before_request
152
  def lazy_init():
153
  """Initialise heavy components before the first real API call."""
154
- if request.path in _LAZY_INIT_EXEMPT_PATHS:
155
- return
156
- if request.path.startswith('/api/output-files/'):
157
- return
158
- if request.path.startswith('/api/models/') and request.path.endswith('/download/progress'):
159
- return
160
  try:
161
  _ensure_components()
162
  except Exception as e:
163
  if request.path.startswith('/api/'):
164
  return jsonify({'error': f'Backend not ready: {e}'}), 503
 
165
  return None
166
 
167
 
@@ -177,51 +129,11 @@ def health_check():
177
  'gpu_name': torch.cuda.get_device_name(0) if torch.cuda.is_available() else None,
178
  }
179
  code = 200 if _components_initialised else 503
180
-
 
181
  return jsonify(status), 200
182
 
183
 
184
- @app.route('/api/environment')
185
- def get_environment():
186
- """Return runtime environment info (Docker vs desktop)."""
187
- is_docker = os.environ.get('FRAGMENTA_DOCKER', '').strip() == '1'
188
- return jsonify({'docker': is_docker})
189
-
190
-
191
- @app.route('/api/hf-token/status', methods=['GET'])
192
- def hf_token_status():
193
- """Check if a valid Hugging Face token is configured."""
194
- try:
195
- from huggingface_hub import HfApi
196
- api = HfApi()
197
- user = api.whoami()
198
- return jsonify({'authenticated': True, 'username': user.get('name', user.get('fullname', 'Unknown'))})
199
- except Exception:
200
- return jsonify({'authenticated': False, 'username': None})
201
-
202
-
203
- @app.route('/api/hf-token', methods=['POST'])
204
- def set_hf_token():
205
- """Set a Hugging Face token (for Docker mode where there's no desktop dialog)."""
206
- try:
207
- data = request.get_json()
208
- token = data.get('token', '').strip()
209
- if not token:
210
- return jsonify({'error': 'Token is required'}), 400
211
-
212
- from huggingface_hub import login, HfApi
213
- login(token=token, add_to_git_credential=False)
214
-
215
- api = HfApi()
216
- user = api.whoami()
217
- username = user.get('name', user.get('fullname', 'Unknown'))
218
- logger.info(f"HF token set successfully for user: {username}")
219
- return jsonify({'success': True, 'username': username})
220
- except Exception as e:
221
- logger.error(f"Error setting HF token: {e}")
222
- return jsonify({'error': f'Invalid token: {str(e)}'}), 400
223
-
224
-
225
  @app.route('/')
226
  def serve_react_app():
227
  return send_from_directory(app.static_folder, 'index.html')
@@ -297,17 +209,17 @@ def process_files():
297
  chunks_preview_data = []
298
  for filename, prompt in prompts_data:
299
  chunks_preview_data.append([
300
- filename,
301
- filename,
302
- prompt,
303
- "original"
304
  ])
305
 
306
-
307
  json_path = Path(config.get_metadata_json_path())
308
  existing_metadata = []
309
 
310
-
311
  if json_path.exists():
312
  try:
313
  with open(json_path, 'r', encoding='utf-8') as f:
@@ -323,34 +235,26 @@ def process_files():
323
  existing_files[filename] = {
324
  "file_name": filename,
325
  "prompt": prompt,
326
- "path": f"data/{filename}"
327
  }
328
 
 
329
  final_metadata = list(existing_files.values())
330
 
331
  with open(json_path, 'w', encoding='utf-8') as f:
332
  import json
333
  json.dump(final_metadata, f, indent=2)
334
 
335
- legacy_json_path = Path(__file__).parent / 'data' / 'metadata.json'
336
- try:
337
- legacy_json_path.parent.mkdir(parents=True, exist_ok=True)
338
- with open(legacy_json_path, 'w', encoding='utf-8') as f:
339
- json.dump(final_metadata, f, indent=2)
340
- except Exception as e:
341
- logger.warning(f"Could not write legacy metadata to {legacy_json_path}: {e}")
342
-
343
-
344
  try:
345
  config.update_dataset_config()
346
- except Exception as e:
347
- logger.warning(f"Could not update dataset config: {e}")
348
 
349
  return jsonify({
350
  'message': f'Files saved successfully! {len(saved_files)} original files saved to data folder',
351
  'saved_files': saved_files,
352
  'processed_count': len(saved_files),
353
- 'chunks_preview': chunks_preview_data,
354
  'data_folder': str(data_dir),
355
  'metadata_json': str(json_path),
356
  'approach': 'original_files_only'
@@ -374,17 +278,11 @@ def start_training():
374
  print(f" - Model Name: {training_config.get('modelName', 'untitled')}")
375
  print(f" - Base Model: {training_config.get('baseModel', 'NOT SET')}")
376
  print(f" - Epochs: {training_config.get('epochs', 'NOT SET')}")
 
377
  print(f" - Batch Size: {training_config.get('batchSize', 'NOT SET')}")
378
  print(f" - Learning Rate: {training_config.get('learningRate', 'NOT SET')}")
379
  print(f" - Save Wrapped Checkpoint: {training_config.get('saveWrappedCheckpoint', False)}")
380
 
381
- try:
382
- _cfg = get_config()
383
- _cfg.update_dataset_config()
384
- print(" Dataset config regenerated successfully")
385
- except Exception as dc_err:
386
- print(f" WARNING: Could not regenerate dataset config: {dc_err}")
387
-
388
  required_fields = ['modelName', 'baseModel']
389
  missing_fields = [field for field in required_fields if field not in training_config]
390
  if missing_fields:
@@ -400,8 +298,11 @@ def start_training():
400
  return jsonify({'error': error_msg}), 400
401
 
402
  if 'epochs' not in training_config:
403
- training_config['epochs'] = 3
404
- print(f" Setting default epochs: 3")
 
 
 
405
  if 'batchSize' not in training_config:
406
  training_config['batchSize'] = 1
407
  print(f" Setting default batch size: 1")
@@ -416,6 +317,7 @@ def start_training():
416
  print(f" - Model Name: {training_config['modelName']}")
417
  print(f" - Base Model: {training_config['baseModel']}")
418
  print(f" - Epochs: {training_config['epochs']}")
 
419
  print(f" - Batch Size: {training_config['batchSize']}")
420
  print(f" - Learning Rate: {training_config['learningRate']}")
421
  print(f" - Save Wrapped Checkpoint: {training_config['saveWrappedCheckpoint']}")
@@ -464,6 +366,7 @@ def generate_audio():
464
  config_file = None
465
  model_file_path = None
466
 
 
467
  if unwrapped_model_path:
468
  model_file_path = Path(unwrapped_model_path)
469
  if not model_file_path.exists():
@@ -478,6 +381,7 @@ def generate_audio():
478
  f"model_path:{model_name}", str(model_file_path))
479
  logger.debug(f"Using model path: {model_file_path}")
480
 
 
481
  if model_file_path:
482
  file_size_gb = model_file_path.stat().st_size / (1024**3)
483
  config_file = "model_config_small.json" if file_size_gb < 2.0 else "model_config.json"
@@ -498,7 +402,7 @@ def generate_audio():
498
  logger.info(f"Starting generation with config: {config_file}")
499
  try:
500
  if determined_model_path and determined_model_path.exists():
501
-
502
  output_path = generator.generate_audio(
503
  prompt,
504
  unwrapped_model_path=unwrapped_model_path if unwrapped_model_path else None,
@@ -507,7 +411,7 @@ def generate_audio():
507
  duration=duration
508
  )
509
  elif model_name in ['stable-audio-open-small', 'stable-audio-open-1.0']:
510
-
511
  model_file_mapping = {
512
  'stable-audio-open-small': 'stable-audio-open-small-model.safetensors',
513
  'stable-audio-open-1.0': 'stable-audio-open-model.safetensors'
@@ -599,7 +503,7 @@ def get_status():
599
  'has_metadata_json': metadata_json.exists(),
600
  'has_custom_metadata': custom_metadata.exists(),
601
  'trained_models': len(list(config.get_path("models_fine_tuned").glob("*"))) if config.get_path("models_fine_tuned").exists() else 0,
602
- 'training': get_training_status() if get_training_status else {'is_training': False, 'progress': 0}
603
  }
604
 
605
  return jsonify(status_response)
@@ -644,8 +548,10 @@ def get_models():
644
  has_checkpoint = len(checkpoint_files) > 0
645
  has_config = len(config_files) > 0
646
 
 
647
  checkpoints = []
648
  for ckpt_file in checkpoint_files:
 
649
  import re
650
  name = ckpt_file.stem
651
  epoch_match = re.search(r'epoch=(\d+)', name)
@@ -653,6 +559,7 @@ def get_models():
653
 
654
  checkpoint_info = {
655
  'name': name,
 
656
  'path': str(ckpt_file.relative_to(config.project_root)),
657
  'size_mb': round(ckpt_file.stat().st_size / (1024 * 1024), 1),
658
  'created': ckpt_file.stat().st_mtime
@@ -665,38 +572,45 @@ def get_models():
665
 
666
  checkpoints.append(checkpoint_info)
667
 
 
668
  checkpoints.sort(key=lambda x: x['created'], reverse=True)
669
 
 
670
  latest_checkpoint = max(checkpoint_files, key=lambda x: x.stat(
671
  ).st_mtime) if checkpoint_files else None
672
  latest_config = max(
673
  config_files, key=lambda x: x.stat().st_mtime) if config_files else None
674
 
 
675
  unwrapped_dir = model_dir / "unwrapped"
676
  unwrapped_models = []
677
  if unwrapped_dir.exists():
678
  for unwrapped_file in unwrapped_dir.glob("*.safetensors"):
679
  unwrapped_models.append({
680
  'name': unwrapped_file.stem,
 
681
  'path': str(unwrapped_file.relative_to(config.project_root)),
682
  'size_mb': round(unwrapped_file.stat().st_size / (1024 * 1024), 1),
683
  'created': unwrapped_file.stat().st_mtime
684
  })
685
 
 
686
  unwrapped_models.sort(
687
  key=lambda x: x['created'], reverse=True)
688
 
689
- base_config_path = "models/config/model_config_small.json"
 
690
 
691
  models.append({
692
  'name': model_dir.name,
 
693
  'path': str(model_dir.relative_to(config.project_root)),
694
  'has_checkpoint': has_checkpoint,
695
  'has_config': has_config,
696
  # Use relative path
697
  'ckpt_path': str(latest_checkpoint.relative_to(config.project_root)) if latest_checkpoint else None,
698
- 'config_path': base_config_path,
699
- 'checkpoints': checkpoints,
700
  'unwrapped_models': unwrapped_models,
701
  'created': model_dir.stat().st_mtime if model_dir.exists() else None
702
  })
@@ -743,43 +657,43 @@ def accept_model_terms(model_id):
743
 
744
  @app.route('/api/models/<model_id>/download', methods=['POST'])
745
  def download_model(model_id):
746
- """Start downloading a model from Hugging Face (async)"""
747
  try:
 
748
  if not model_manager.is_terms_accepted(model_id):
749
  return jsonify({'error': 'Terms not accepted for this model'}), 400
750
 
751
- if model_id in _download_progress and _download_progress[model_id].get('status') == 'downloading':
752
- return jsonify({'success': True, 'message': 'Download already in progress'})
753
-
754
- _download_progress[model_id] = {'percent': 0, 'message': 'Starting download...', 'status': 'downloading'}
755
-
756
- def progress_callback(percent, message):
757
- _download_progress[model_id] = {'percent': percent, 'message': message, 'status': 'downloading'}
758
-
759
- def run_download():
760
- try:
761
- success = model_manager.download_model(model_id, progress_callback=progress_callback)
762
- if success:
763
- _download_progress[model_id] = {'percent': 100, 'message': 'Download complete', 'status': 'done'}
764
- else:
765
- msg = _download_progress.get(model_id, {}).get('message', 'Download failed')
766
- _download_progress[model_id] = {'percent': 0, 'message': msg, 'status': 'error'}
767
- except Exception as e:
768
- _download_progress[model_id] = {'percent': 0, 'message': str(e), 'status': 'error'}
769
-
770
- t = threading.Thread(target=run_download, daemon=True)
771
- t.start()
772
-
773
- return jsonify({'success': True, 'message': f'Download started for {model_id}'})
774
  except Exception as e:
775
  return jsonify({'error': str(e)}), 500
776
 
777
 
778
- @app.route('/api/models/<model_id>/download/progress', methods=['GET'])
779
- def download_progress(model_id):
780
- """Get the download progress for a model"""
781
- progress = _download_progress.get(model_id, {'percent': 0, 'message': 'No download in progress', 'status': 'idle'})
782
- return jsonify(progress)
 
 
 
 
 
 
 
 
 
 
 
 
 
783
 
784
 
785
  @app.route('/api/base-models/status', methods=['GET'])
@@ -789,34 +703,31 @@ def get_base_models_status():
789
  import os
790
  from pathlib import Path
791
 
792
- # Use absolute path based on project root (reliable in Docker)
793
- project_root = Path(__file__).parent.parent.parent
794
- pretrained_dir = project_root / 'models' / 'pretrained'
795
-
796
  base_models = {
797
  'stable-audio-open-1.0': {
798
  'name': 'Stable Audio Open 1.0',
799
- 'path': str(pretrained_dir),
800
- 'file': 'stable-audio-open-model.safetensors',
801
  'downloaded': False
802
  },
803
  'stable-audio-open-small': {
804
  'name': 'Stable Audio Open Small',
805
- 'path': str(pretrained_dir),
806
- 'file': 'stable-audio-open-small-model.safetensors',
807
  'downloaded': False
808
  }
809
  }
810
 
 
811
  for model_id, info in base_models.items():
812
  model_dir = Path(info['path'])
813
  model_file = model_dir / info['file']
814
 
815
- logger.info(f"[base-models/status] Checking {model_id}: {model_file} exists={model_file.exists()}")
816
-
817
  if model_file.exists() and model_file.is_file():
818
  info['downloaded'] = True
819
  else:
 
820
  old_path = model_dir / model_id
821
  if old_path.exists() and old_path.is_dir():
822
  has_files = any([
@@ -864,6 +775,7 @@ def start_fresh():
864
  data_dir = config.get_path("data")
865
  config_dir = config.get_path("models_config")
866
 
 
867
  data_files_deleted = 0
868
  if data_dir.exists():
869
  for file_path in data_dir.glob("*"):
@@ -871,6 +783,7 @@ def start_fresh():
871
  file_path.unlink()
872
  data_files_deleted += 1
873
 
 
874
  config_files_deleted = 0
875
  if config_dir.exists():
876
  for file_path in config_dir.glob("custom_metadata.py"):
@@ -878,6 +791,7 @@ def start_fresh():
878
  file_path.unlink()
879
  config_files_deleted += 1
880
 
 
881
  data_dir.mkdir(exist_ok=True, parents=True)
882
 
883
  return jsonify({
@@ -902,24 +816,28 @@ def unwrap_model():
902
  if not model_config or not ckpt_path:
903
  return jsonify({'error': 'model_config and ckpt_path are required'}), 400
904
 
 
905
  import subprocess
906
  from pathlib import Path
907
 
 
908
  config = get_config()
909
  repo_root = config.project_root
910
 
 
911
  model_config_path = repo_root / \
912
  model_config if not Path(
913
  model_config).is_absolute() else Path(model_config)
914
  ckpt_path_resolved = repo_root / \
915
  ckpt_path if not Path(ckpt_path).is_absolute() else Path(ckpt_path)
916
 
 
917
  if not model_config_path.exists():
918
  return jsonify({'error': f'Model config not found: {model_config_path}'}), 400
919
  if not ckpt_path_resolved.exists():
920
  return jsonify({'error': f'Checkpoint not found: {ckpt_path_resolved}'}), 400
921
 
922
-
923
  model_dir = ckpt_path_resolved.parent
924
  unwrapped_dir = model_dir / "unwrapped"
925
  unwrapped_dir.mkdir(exist_ok=True)
@@ -1300,24 +1218,14 @@ _memory_warning_interval = 30 # seconds
1300
 
1301
  @app.route('/api/open-output-folder', methods=['POST'])
1302
  def open_output_folder():
1303
- """Open the output folder in the system file explorer.
1304
- In Docker mode, returns a helpful message instead of trying xdg-open."""
1305
  try:
1306
- output_path = Path("output")
1307
- output_path.mkdir(exist_ok=True)
1308
-
1309
- is_docker = os.environ.get('FRAGMENTA_DOCKER', '').strip() == '1'
1310
- if is_docker:
1311
- return jsonify({
1312
- "success": True,
1313
- "docker": True,
1314
- "path": str(output_path.absolute()),
1315
- "message": "In Docker mode, access output files via the mounted volume on your host (./output/) or use the file browser in the UI."
1316
- })
1317
-
1318
  import subprocess
1319
  import platform
1320
-
 
 
 
1321
  system = platform.system()
1322
  if system == "Windows":
1323
  subprocess.run(["explorer", str(output_path.absolute())])
@@ -1325,61 +1233,41 @@ def open_output_folder():
1325
  subprocess.run(["open", str(output_path.absolute())])
1326
  else: # Linux
1327
  subprocess.run(["xdg-open", str(output_path.absolute())])
1328
-
1329
  return jsonify({"success": True, "message": "Output folder opened"})
1330
  except Exception as e:
1331
  logger.error(f"Error opening output folder: {e}")
1332
  return jsonify({"success": False, "error": str(e)}), 500
1333
 
1334
-
1335
- @app.route('/api/output-files', methods=['GET'])
1336
- def list_output_files():
1337
- """List files in the output directory (useful for Docker users)."""
1338
  try:
1339
- output_path = Path("output")
1340
- output_path.mkdir(exist_ok=True)
1341
 
1342
- files = []
1343
- for f in sorted(output_path.iterdir(), key=lambda p: p.stat().st_mtime, reverse=True):
1344
- if f.is_file():
1345
- stat = f.stat()
1346
- files.append({
1347
- 'name': f.name,
1348
- 'size': stat.st_size,
1349
- 'modified': stat.st_mtime
1350
- })
1351
- return jsonify({'files': files})
1352
- except Exception as e:
1353
- logger.error(f"Error listing output files: {e}")
1354
- return jsonify({'error': str(e)}), 500
1355
 
 
 
 
 
1356
 
1357
- @app.route('/api/output-files/<path:filename>', methods=['GET'])
1358
- def download_output_file(filename):
1359
- """Download a specific file from the output directory."""
1360
- try:
1361
- from werkzeug.utils import secure_filename as _secure
1362
- safe_name = _secure(filename)
1363
- output_path = Path("output")
1364
- file_path = output_path / safe_name
1365
- if not file_path.exists() or not file_path.is_file():
1366
- return jsonify({'error': 'File not found'}), 404
1367
- return send_file(str(file_path.absolute()), as_attachment=True, download_name=safe_name)
1368
- except Exception as e:
1369
- logger.error(f"Error downloading output file: {e}")
1370
- return jsonify({'error': str(e)}), 500
1371
 
1372
- @app.route('/api/open-documentation', methods=['POST'])
1373
- def open_documentation():
1374
- """Open the documentation URL in the default browser"""
1375
- try:
1376
- import webbrowser
1377
-
1378
- # TODO: Replace with actual documentation URL
1379
- documentation_url = "https://github.com/your-repo/fragmenta-docs"
1380
- webbrowser.open(documentation_url)
1381
-
1382
- return jsonify({"success": True, "message": "Documentation opened"})
1383
  except Exception as e:
1384
  logger.error(f"Error opening documentation: {e}")
1385
  return jsonify({"success": False, "error": str(e)}), 500
@@ -1446,10 +1334,31 @@ def get_license_info():
1446
  def get_models_status():
1447
  """Check if required models exist and if auth dialog should be shown"""
1448
  try:
1449
- from app.core.hf_auth_dialog import check_required_models_exist, should_show_auth_dialog
1450
-
1451
- models_exist, models_message = check_required_models_exist()
1452
- should_show, auth_reason = should_show_auth_dialog()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1453
 
1454
  return jsonify({
1455
  "models_exist": models_exist,
@@ -1575,6 +1484,267 @@ def get_gpu_memory_status():
1575
  return jsonify({'error': str(e)}), 500
1576
 
1577
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1578
  @app.route('/shutdown', methods=['POST'])
1579
  def shutdown():
1580
  """Shutdown the Flask server gracefully"""
@@ -1591,4 +1761,7 @@ def shutdown():
1591
 
1592
 
1593
  if __name__ == '__main__':
1594
- app.run(debug=True, port=5001)
 
 
 
 
2
  from utils.exceptions import ModelNotFoundError, ValidationError, GenerationError
3
  from utils.api_responses import APIResponse, handle_api_error
4
  from utils.logger import setup_logging, get_logger
5
+ from app.core.generation.audio_generator import AudioGenerator
6
+ from app.core.training.lora_trainer import start_training as start_training_func, get_training_status, stop_training
7
+ from app.backend.data.simple_audio_processor import SimpleAudioProcessor
8
  from app.core.config import get_config
9
  from flask import Flask, request, jsonify, send_file, send_from_directory
10
  from flask_cors import CORS
 
17
  import logging
18
  from werkzeug.serving import WSGIRequestHandler
19
 
 
 
 
 
 
 
 
20
  sys.path.append(os.path.abspath(
21
  os.path.join(os.path.dirname(__file__), '../../')))
22
 
 
58
 
59
  DEBUG_MODE = os.environ.get('FRAGMENTA_DEBUG', 'false').lower() == 'true'
60
 
61
+ # ---------------------------------------------------------------------------
62
+ # Lazy-initialised backend components
63
+ # ---------------------------------------------------------------------------
64
+ # These are initialised on first real API request (not at import time) so that
65
+ # the Flask server always starts — even when model files or heavy deps are
66
+ # temporarily unavailable. The /api/health endpoint works unconditionally.
67
+ # ---------------------------------------------------------------------------
68
  config = None
69
  audio_processor = None
70
  generator = None
 
74
 
75
 
76
  def _ensure_components():
77
+ """Initialise backend components on first use. Thread-safe."""
 
 
 
 
78
  global config, audio_processor, generator, model_manager
79
  global _components_initialised, _init_error
80
 
81
  if _components_initialised:
82
  return
83
+ if _init_error:
84
+ raise RuntimeError(f"Backend failed to initialise earlier: {_init_error}")
85
 
86
  try:
87
  logger.info("Initializing Backend API components (lazy)…")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
88
  config = get_config()
89
  audio_processor = SimpleAudioProcessor(
90
  model_config_path=config.get_path("models_config") / "model_config.json"
 
95
  model_manager = ModelManager(config)
96
 
97
  _components_initialised = True
 
98
  logger.info("Backend components initialized successfully")
99
 
100
  except Exception as e:
 
102
  logger.error(f"Failed to initialize backend components: {e}")
103
  raise
104
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
105
 
106
  @app.before_request
107
  def lazy_init():
108
  """Initialise heavy components before the first real API call."""
109
+ if request.path == '/api/health':
110
+ return # health endpoint must always work
 
 
 
 
111
  try:
112
  _ensure_components()
113
  except Exception as e:
114
  if request.path.startswith('/api/'):
115
  return jsonify({'error': f'Backend not ready: {e}'}), 503
116
+ # Static file / React routes — let them through even if init fails
117
  return None
118
 
119
 
 
129
  'gpu_name': torch.cuda.get_device_name(0) if torch.cuda.is_available() else None,
130
  }
131
  code = 200 if _components_initialised else 503
132
+ # Return 200 even in degraded mode so Docker HEALTHCHECK doesn't kill
133
+ # the container before components finish loading
134
  return jsonify(status), 200
135
 
136
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
137
  @app.route('/')
138
  def serve_react_app():
139
  return send_from_directory(app.static_folder, 'index.html')
 
209
  chunks_preview_data = []
210
  for filename, prompt in prompts_data:
211
  chunks_preview_data.append([
212
+ filename, # Original filename (not chunked)
213
+ filename, # Source file
214
+ prompt, # User's prompt
215
+ "original" # Not chunked
216
  ])
217
 
218
+ # Do not overwrite the metadata! keeps dataset creation more sustainable
219
  json_path = Path(config.get_metadata_json_path())
220
  existing_metadata = []
221
 
222
+ # Load existing metadata if file exists
223
  if json_path.exists():
224
  try:
225
  with open(json_path, 'r', encoding='utf-8') as f:
 
235
  existing_files[filename] = {
236
  "file_name": filename,
237
  "prompt": prompt,
238
+ "path": f"app/backend/data/{filename}"
239
  }
240
 
241
+ # Convert back to list and save
242
  final_metadata = list(existing_files.values())
243
 
244
  with open(json_path, 'w', encoding='utf-8') as f:
245
  import json
246
  json.dump(final_metadata, f, indent=2)
247
 
 
 
 
 
 
 
 
 
 
248
  try:
249
  config.update_dataset_config()
250
+ except Exception as exc:
251
+ print(f"Warning: failed to refresh dataset-config.json: {exc}")
252
 
253
  return jsonify({
254
  'message': f'Files saved successfully! {len(saved_files)} original files saved to data folder',
255
  'saved_files': saved_files,
256
  'processed_count': len(saved_files),
257
+ 'chunks_preview': chunks_preview_data, # Show all files (no chunking)
258
  'data_folder': str(data_dir),
259
  'metadata_json': str(json_path),
260
  'approach': 'original_files_only'
 
278
  print(f" - Model Name: {training_config.get('modelName', 'untitled')}")
279
  print(f" - Base Model: {training_config.get('baseModel', 'NOT SET')}")
280
  print(f" - Epochs: {training_config.get('epochs', 'NOT SET')}")
281
+ print(f" - Checkpoint Steps: {training_config.get('checkpointSteps', 'NOT SET')}")
282
  print(f" - Batch Size: {training_config.get('batchSize', 'NOT SET')}")
283
  print(f" - Learning Rate: {training_config.get('learningRate', 'NOT SET')}")
284
  print(f" - Save Wrapped Checkpoint: {training_config.get('saveWrappedCheckpoint', False)}")
285
 
 
 
 
 
 
 
 
286
  required_fields = ['modelName', 'baseModel']
287
  missing_fields = [field for field in required_fields if field not in training_config]
288
  if missing_fields:
 
298
  return jsonify({'error': error_msg}), 400
299
 
300
  if 'epochs' not in training_config:
301
+ training_config['epochs'] = 30
302
+ print(f" Setting default epochs: 30")
303
+ if 'checkpointSteps' not in training_config:
304
+ training_config['checkpointSteps'] = 50
305
+ print(f" Setting default checkpointSteps: 50")
306
  if 'batchSize' not in training_config:
307
  training_config['batchSize'] = 1
308
  print(f" Setting default batch size: 1")
 
317
  print(f" - Model Name: {training_config['modelName']}")
318
  print(f" - Base Model: {training_config['baseModel']}")
319
  print(f" - Epochs: {training_config['epochs']}")
320
+ print(f" - Checkpoint Steps: {training_config['checkpointSteps']}")
321
  print(f" - Batch Size: {training_config['batchSize']}")
322
  print(f" - Learning Rate: {training_config['learningRate']}")
323
  print(f" - Save Wrapped Checkpoint: {training_config['saveWrappedCheckpoint']}")
 
366
  config_file = None
367
  model_file_path = None
368
 
369
+ # Priority: unwrapped_model_path > model_path > base model
370
  if unwrapped_model_path:
371
  model_file_path = Path(unwrapped_model_path)
372
  if not model_file_path.exists():
 
381
  f"model_path:{model_name}", str(model_file_path))
382
  logger.debug(f"Using model path: {model_file_path}")
383
 
384
+ # Determine config based on file size or model name
385
  if model_file_path:
386
  file_size_gb = model_file_path.stat().st_size / (1024**3)
387
  config_file = "model_config_small.json" if file_size_gb < 2.0 else "model_config.json"
 
402
  logger.info(f"Starting generation with config: {config_file}")
403
  try:
404
  if determined_model_path and determined_model_path.exists():
405
+ # Use the determined model path
406
  output_path = generator.generate_audio(
407
  prompt,
408
  unwrapped_model_path=unwrapped_model_path if unwrapped_model_path else None,
 
411
  duration=duration
412
  )
413
  elif model_name in ['stable-audio-open-small', 'stable-audio-open-1.0']:
414
+ # Handle base models
415
  model_file_mapping = {
416
  'stable-audio-open-small': 'stable-audio-open-small-model.safetensors',
417
  'stable-audio-open-1.0': 'stable-audio-open-model.safetensors'
 
503
  'has_metadata_json': metadata_json.exists(),
504
  'has_custom_metadata': custom_metadata.exists(),
505
  'trained_models': len(list(config.get_path("models_fine_tuned").glob("*"))) if config.get_path("models_fine_tuned").exists() else 0,
506
+ 'training': get_training_status()
507
  }
508
 
509
  return jsonify(status_response)
 
548
  has_checkpoint = len(checkpoint_files) > 0
549
  has_config = len(config_files) > 0
550
 
551
+ # Create detailed checkpoint information
552
  checkpoints = []
553
  for ckpt_file in checkpoint_files:
554
+ # Extract epoch and step from filename if possible
555
  import re
556
  name = ckpt_file.stem
557
  epoch_match = re.search(r'epoch=(\d+)', name)
 
559
 
560
  checkpoint_info = {
561
  'name': name,
562
+ # Use relative path
563
  'path': str(ckpt_file.relative_to(config.project_root)),
564
  'size_mb': round(ckpt_file.stat().st_size / (1024 * 1024), 1),
565
  'created': ckpt_file.stat().st_mtime
 
572
 
573
  checkpoints.append(checkpoint_info)
574
 
575
+ # Sort checkpoints by creation time (newest first)
576
  checkpoints.sort(key=lambda x: x['created'], reverse=True)
577
 
578
+ # Get the latest checkpoint and config files
579
  latest_checkpoint = max(checkpoint_files, key=lambda x: x.stat(
580
  ).st_mtime) if checkpoint_files else None
581
  latest_config = max(
582
  config_files, key=lambda x: x.stat().st_mtime) if config_files else None
583
 
584
+ # Check for unwrapped models
585
  unwrapped_dir = model_dir / "unwrapped"
586
  unwrapped_models = []
587
  if unwrapped_dir.exists():
588
  for unwrapped_file in unwrapped_dir.glob("*.safetensors"):
589
  unwrapped_models.append({
590
  'name': unwrapped_file.stem,
591
+ # Use relative path
592
  'path': str(unwrapped_file.relative_to(config.project_root)),
593
  'size_mb': round(unwrapped_file.stat().st_size / (1024 * 1024), 1),
594
  'created': unwrapped_file.stat().st_mtime
595
  })
596
 
597
+ # Sort unwrapped models by creation time (newest first)
598
  unwrapped_models.sort(
599
  key=lambda x: x['created'], reverse=True)
600
 
601
+ # For fine-tuned models, use the base model's config
602
+ base_config_path = "models/config/model_config_small.json" # Use relative path
603
 
604
  models.append({
605
  'name': model_dir.name,
606
+ # Use relative path
607
  'path': str(model_dir.relative_to(config.project_root)),
608
  'has_checkpoint': has_checkpoint,
609
  'has_config': has_config,
610
  # Use relative path
611
  'ckpt_path': str(latest_checkpoint.relative_to(config.project_root)) if latest_checkpoint else None,
612
+ 'config_path': base_config_path, # Use base model config for unwrapping
613
+ 'checkpoints': checkpoints, # Detailed checkpoint list
614
  'unwrapped_models': unwrapped_models,
615
  'created': model_dir.stat().st_mtime if model_dir.exists() else None
616
  })
 
657
 
658
  @app.route('/api/models/<model_id>/download', methods=['POST'])
659
  def download_model(model_id):
660
+ """Download a model from Hugging Face"""
661
  try:
662
+ # Check if terms are accepted
663
  if not model_manager.is_terms_accepted(model_id):
664
  return jsonify({'error': 'Terms not accepted for this model'}), 400
665
 
666
+ # Start download
667
+ success = model_manager.download_model(model_id)
668
+ if success:
669
+ return jsonify({
670
+ 'success': True,
671
+ 'message': f'Model {model_id} downloaded successfully'
672
+ })
673
+ else:
674
+ return jsonify({'error': f'Failed to download {model_id}'}), 500
 
 
 
 
 
 
 
 
 
 
 
 
 
 
675
  except Exception as e:
676
  return jsonify({'error': str(e)}), 500
677
 
678
 
679
+ @app.route('/api/hf-login', methods=['POST'])
680
+ def hf_login():
681
+ """Login to Hugging Face with a token"""
682
+ try:
683
+ data = request.json
684
+ token = data.get('token')
685
+ if not token:
686
+ return jsonify({'error': 'Token is required'}), 400
687
+
688
+ import huggingface_hub
689
+ try:
690
+ huggingface_hub.login(token=token, add_to_git_credential=False)
691
+ user_info = huggingface_hub.whoami(token=token)
692
+ return jsonify({'success': True, 'user': user_info.get('name', 'User')})
693
+ except Exception as e:
694
+ return jsonify({'error': f'Invalid token or connection error: {str(e)}'}), 401
695
+ except Exception as e:
696
+ return jsonify({'error': str(e)}), 500
697
 
698
 
699
  @app.route('/api/base-models/status', methods=['GET'])
 
703
  import os
704
  from pathlib import Path
705
 
 
 
 
 
706
  base_models = {
707
  'stable-audio-open-1.0': {
708
  'name': 'Stable Audio Open 1.0',
709
+ 'path': 'models/pretrained', # Updated to correct path
710
+ 'file': 'stable-audio-open-model.safetensors', # Specific file to check
711
  'downloaded': False
712
  },
713
  'stable-audio-open-small': {
714
  'name': 'Stable Audio Open Small',
715
+ 'path': 'models/pretrained', # Updated to correct path
716
+ 'file': 'stable-audio-open-small-model.safetensors', # Specific file to check
717
  'downloaded': False
718
  }
719
  }
720
 
721
+ # Check if models are actually downloaded by looking for specific files
722
  for model_id, info in base_models.items():
723
  model_dir = Path(info['path'])
724
  model_file = model_dir / info['file']
725
 
726
+ # Check if the specific model file exists
 
727
  if model_file.exists() and model_file.is_file():
728
  info['downloaded'] = True
729
  else:
730
+ # Fallback: check subdirectory structure (old format)
731
  old_path = model_dir / model_id
732
  if old_path.exists() and old_path.is_dir():
733
  has_files = any([
 
775
  data_dir = config.get_path("data")
776
  config_dir = config.get_path("models_config")
777
 
778
+ # Delete all data files
779
  data_files_deleted = 0
780
  if data_dir.exists():
781
  for file_path in data_dir.glob("*"):
 
783
  file_path.unlink()
784
  data_files_deleted += 1
785
 
786
+ # Delete config metadata files (but keep the model configs)
787
  config_files_deleted = 0
788
  if config_dir.exists():
789
  for file_path in config_dir.glob("custom_metadata.py"):
 
791
  file_path.unlink()
792
  config_files_deleted += 1
793
 
794
+ # Recreate empty data directory
795
  data_dir.mkdir(exist_ok=True, parents=True)
796
 
797
  return jsonify({
 
816
  if not model_config or not ckpt_path:
817
  return jsonify({'error': 'model_config and ckpt_path are required'}), 400
818
 
819
+ # Use the stable-audio-tools unwrap_model.py script directly for individual checkpoints
820
  import subprocess
821
  from pathlib import Path
822
 
823
+ # Get config to resolve relative paths
824
  config = get_config()
825
  repo_root = config.project_root
826
 
827
+ # Resolve paths relative to project root
828
  model_config_path = repo_root / \
829
  model_config if not Path(
830
  model_config).is_absolute() else Path(model_config)
831
  ckpt_path_resolved = repo_root / \
832
  ckpt_path if not Path(ckpt_path).is_absolute() else Path(ckpt_path)
833
 
834
+ # Validate paths exist
835
  if not model_config_path.exists():
836
  return jsonify({'error': f'Model config not found: {model_config_path}'}), 400
837
  if not ckpt_path_resolved.exists():
838
  return jsonify({'error': f'Checkpoint not found: {ckpt_path_resolved}'}), 400
839
 
840
+ # Get the model directory and create unwrapped subdirectory
841
  model_dir = ckpt_path_resolved.parent
842
  unwrapped_dir = model_dir / "unwrapped"
843
  unwrapped_dir.mkdir(exist_ok=True)
 
1218
 
1219
  @app.route('/api/open-output-folder', methods=['POST'])
1220
  def open_output_folder():
1221
+ """Open the output folder in the system file explorer"""
 
1222
  try:
 
 
 
 
 
 
 
 
 
 
 
 
1223
  import subprocess
1224
  import platform
1225
+
1226
+ output_path = Path("output")
1227
+ output_path.mkdir(exist_ok=True)
1228
+
1229
  system = platform.system()
1230
  if system == "Windows":
1231
  subprocess.run(["explorer", str(output_path.absolute())])
 
1233
  subprocess.run(["open", str(output_path.absolute())])
1234
  else: # Linux
1235
  subprocess.run(["xdg-open", str(output_path.absolute())])
1236
+
1237
  return jsonify({"success": True, "message": "Output folder opened"})
1238
  except Exception as e:
1239
  logger.error(f"Error opening output folder: {e}")
1240
  return jsonify({"success": False, "error": str(e)}), 500
1241
 
1242
+ @app.route('/api/open-documentation', methods=['POST'])
1243
+ def open_documentation():
1244
+ """Open selected public Fragmenta links in the system browser."""
 
1245
  try:
1246
+ import webbrowser
 
1247
 
1248
+ payload = request.get_json(silent=True) or {}
1249
+ doc_key = payload.get('doc_key', 'about')
 
 
 
 
 
 
 
 
 
 
 
1250
 
1251
+ docs_map = {
1252
+ 'about': 'https://www.misaghazimi.com/fragmenta',
1253
+ 'documentation': 'https://github.com/MAz-Codes/Fragmenta',
1254
+ }
1255
 
1256
+ target_url = docs_map.get(doc_key)
1257
+ if not target_url:
1258
+ return jsonify({
1259
+ "success": False,
1260
+ "error": f"Unsupported documentation target: {doc_key}"
1261
+ }), 400
 
 
 
 
 
 
 
 
1262
 
1263
+ webbrowser.open(target_url)
1264
+
1265
+ return jsonify({
1266
+ "success": True,
1267
+ "message": f"Opened {doc_key}",
1268
+ "doc_key": doc_key,
1269
+ "url": target_url,
1270
+ })
 
 
 
1271
  except Exception as e:
1272
  logger.error(f"Error opening documentation: {e}")
1273
  return jsonify({"success": False, "error": str(e)}), 500
 
1334
  def get_models_status():
1335
  """Check if required models exist and if auth dialog should be shown"""
1336
  try:
1337
+ required_models = ['stable-audio-open-small', 'stable-audio-open-1.0']
1338
+ downloaded_models = [
1339
+ model_id for model_id in required_models if model_manager.is_model_downloaded(model_id)
1340
+ ]
1341
+ models_exist = len(downloaded_models) > 0
1342
+ models_message = (
1343
+ "Required base models are available."
1344
+ if models_exist
1345
+ else "No required base model is downloaded yet."
1346
+ )
1347
+
1348
+ hf_authenticated = False
1349
+ try:
1350
+ from huggingface_hub import HfApi
1351
+ HfApi().whoami()
1352
+ hf_authenticated = True
1353
+ except Exception:
1354
+ hf_authenticated = False
1355
+
1356
+ should_show = (not models_exist) and (not hf_authenticated)
1357
+ auth_reason = (
1358
+ "Hugging Face authentication is required to download gated models."
1359
+ if should_show
1360
+ else "Authentication already available or models already downloaded."
1361
+ )
1362
 
1363
  return jsonify({
1364
  "models_exist": models_exist,
 
1484
  return jsonify({'error': str(e)}), 500
1485
 
1486
 
1487
+ # ---------------------------------------------------------------------------
1488
+ # Bulk auto-annotation
1489
+ # ---------------------------------------------------------------------------
1490
+ _annotate_job_lock = threading.Lock()
1491
+ _annotate_job = {
1492
+ 'state': 'idle', # idle | running | done | error
1493
+ 'current': 0,
1494
+ 'total': 0,
1495
+ 'current_file': '',
1496
+ 'tier': None,
1497
+ 'folder': None,
1498
+ 'results': [],
1499
+ 'error': None,
1500
+ }
1501
+ _clap_download_job = {
1502
+ 'state': 'idle', # idle | running | done | error
1503
+ 'message': '',
1504
+ 'error': None,
1505
+ }
1506
+
1507
+
1508
+ def _annotator_labels_path():
1509
+ return Path(get_config().project_root) / 'config' / 'annotator_labels.json'
1510
+
1511
+
1512
+ def _clap_ckpt_path():
1513
+ from app.backend.data.auto_annotator import clap_checkpoint_path
1514
+ return clap_checkpoint_path(get_config().get_path('models_pretrained'))
1515
+
1516
+
1517
+ @app.route('/api/pick-folder', methods=['POST'])
1518
+ def pick_folder():
1519
+ """Open a native folder-picker dialog on the host and return the chosen path."""
1520
+ import subprocess
1521
+ import shutil as _shutil
1522
+
1523
+ payload = request.json or {}
1524
+ start_dir = payload.get('start_dir') or str(Path.home())
1525
+
1526
+ def _try(cmd):
1527
+ try:
1528
+ out = subprocess.run(cmd, capture_output=True, text=True, timeout=300)
1529
+ if out.returncode == 0:
1530
+ path = out.stdout.strip()
1531
+ if path:
1532
+ return path
1533
+ except Exception as exc:
1534
+ logger.debug("folder dialog attempt failed (%s): %s", cmd[0], exc)
1535
+ return None
1536
+
1537
+ chosen = None
1538
+
1539
+ if sys.platform.startswith('linux'):
1540
+ if _shutil.which('zenity'):
1541
+ chosen = _try(['zenity', '--file-selection', '--directory',
1542
+ f'--filename={start_dir}/', '--title=Choose audio folder'])
1543
+ if not chosen and _shutil.which('kdialog'):
1544
+ chosen = _try(['kdialog', '--getexistingdirectory', start_dir])
1545
+ if not chosen:
1546
+ # fall back to system python3's tkinter (the venv may not have tk)
1547
+ script = (
1548
+ "import tkinter as tk; from tkinter import filedialog; "
1549
+ "r = tk.Tk(); r.withdraw(); "
1550
+ f"p = filedialog.askdirectory(initialdir={start_dir!r}, title='Choose audio folder'); "
1551
+ "print(p or '')"
1552
+ )
1553
+ chosen = _try(['python3', '-c', script])
1554
+ elif sys.platform == 'darwin':
1555
+ chosen = _try(['osascript', '-e',
1556
+ f'POSIX path of (choose folder with prompt "Choose audio folder" default location POSIX file "{start_dir}")'])
1557
+ elif sys.platform == 'win32':
1558
+ ps = (
1559
+ "Add-Type -AssemblyName System.Windows.Forms; "
1560
+ "$d = New-Object System.Windows.Forms.FolderBrowserDialog; "
1561
+ f"$d.SelectedPath = '{start_dir}'; "
1562
+ "if ($d.ShowDialog() -eq 'OK') {{ Write-Output $d.SelectedPath }}"
1563
+ )
1564
+ chosen = _try(['powershell', '-NoProfile', '-Command', ps])
1565
+
1566
+ if not chosen:
1567
+ return jsonify({'path': None, 'cancelled': True})
1568
+ return jsonify({'path': chosen})
1569
+
1570
+
1571
+ @app.route('/api/bulk-annotate/status', methods=['GET'])
1572
+ def bulk_annotate_status():
1573
+ from app.backend.data.auto_annotator import clap_checkpoint_available
1574
+ with _annotate_job_lock:
1575
+ snapshot = {k: v for k, v in _annotate_job.items() if k != 'results'}
1576
+ snapshot['result_count'] = len(_annotate_job['results'])
1577
+ snapshot['clap_available'] = clap_checkpoint_available(get_config().get_path('models_pretrained'))
1578
+ snapshot['clap_download'] = dict(_clap_download_job)
1579
+ return jsonify(snapshot)
1580
+
1581
+
1582
+ @app.route('/api/bulk-annotate/results', methods=['GET'])
1583
+ def bulk_annotate_results():
1584
+ with _annotate_job_lock:
1585
+ return jsonify({'results': list(_annotate_job['results']), 'state': _annotate_job['state']})
1586
+
1587
+
1588
+ @app.route('/api/bulk-annotate', methods=['POST'])
1589
+ def bulk_annotate():
1590
+ payload = request.json or {}
1591
+ folder = payload.get('folder_path', '').strip()
1592
+ tier = payload.get('tier', 'basic')
1593
+ if tier not in ('basic', 'rich'):
1594
+ return jsonify({'error': f"Invalid tier: {tier}"}), 400
1595
+ if not folder:
1596
+ return jsonify({'error': 'folder_path is required'}), 400
1597
+
1598
+ folder_path = Path(folder).expanduser()
1599
+ if not folder_path.exists() or not folder_path.is_dir():
1600
+ return jsonify({'error': f'Folder not found: {folder_path}'}), 400
1601
+
1602
+ from app.backend.data.auto_annotator import (
1603
+ annotate_folder, load_label_sets, clap_checkpoint_available,
1604
+ )
1605
+
1606
+ if tier == 'rich' and not clap_checkpoint_available(get_config().get_path('models_pretrained')):
1607
+ return jsonify({'error': 'CLAP checkpoint not downloaded yet.'}), 409
1608
+
1609
+ with _annotate_job_lock:
1610
+ if _annotate_job['state'] == 'running':
1611
+ return jsonify({'error': 'An annotation job is already running.'}), 409
1612
+ _annotate_job.update({
1613
+ 'state': 'running', 'current': 0, 'total': 0, 'current_file': '',
1614
+ 'tier': tier, 'folder': str(folder_path), 'results': [], 'error': None,
1615
+ })
1616
+
1617
+ labels = load_label_sets(_annotator_labels_path())
1618
+
1619
+ def progress_cb(i, total, name):
1620
+ with _annotate_job_lock:
1621
+ _annotate_job['current'] = i
1622
+ _annotate_job['total'] = total
1623
+ _annotate_job['current_file'] = name
1624
+
1625
+ def runner():
1626
+ try:
1627
+ results = annotate_folder(
1628
+ folder_path, tier=tier, label_sets=labels,
1629
+ clap_ckpt_path=_clap_ckpt_path() if tier == 'rich' else None,
1630
+ progress_cb=progress_cb,
1631
+ )
1632
+ with _annotate_job_lock:
1633
+ _annotate_job['results'] = results
1634
+ _annotate_job['state'] = 'done'
1635
+ except Exception as exc:
1636
+ logger.exception("Bulk annotation failed")
1637
+ with _annotate_job_lock:
1638
+ _annotate_job['state'] = 'error'
1639
+ _annotate_job['error'] = str(exc)
1640
+
1641
+ threading.Thread(target=runner, daemon=True).start()
1642
+ return jsonify({'message': 'Annotation started', 'tier': tier, 'folder': str(folder_path)})
1643
+
1644
+
1645
+ @app.route('/api/bulk-annotate/commit', methods=['POST'])
1646
+ def bulk_annotate_commit():
1647
+ """Merge user-reviewed annotation results into metadata.json.
1648
+
1649
+ Body: { entries: [{ file_name, prompt, path }, ...], copy_files: bool }
1650
+ """
1651
+ payload = request.json or {}
1652
+ entries = payload.get('entries') or []
1653
+ copy_files = bool(payload.get('copy_files', True))
1654
+ if not entries:
1655
+ return jsonify({'error': 'No entries to commit.'}), 400
1656
+
1657
+ config = get_config()
1658
+ data_dir = config.get_path('data')
1659
+ data_dir.mkdir(exist_ok=True, parents=True)
1660
+
1661
+ json_path = Path(config.get_metadata_json_path())
1662
+ existing_metadata = []
1663
+ if json_path.exists():
1664
+ try:
1665
+ with open(json_path, 'r', encoding='utf-8') as f:
1666
+ existing_metadata = json.load(f)
1667
+ except Exception as exc:
1668
+ logger.warning("Could not load existing metadata: %s", exc)
1669
+ existing_metadata = []
1670
+ existing_files = {item['file_name']: item for item in existing_metadata}
1671
+
1672
+ import shutil
1673
+ committed = 0
1674
+ for entry in entries:
1675
+ file_name = entry.get('file_name')
1676
+ prompt = (entry.get('prompt') or '').strip()
1677
+ src_path = entry.get('path')
1678
+ if not file_name or not prompt or not src_path:
1679
+ continue
1680
+
1681
+ src = Path(src_path)
1682
+ if copy_files and src.exists() and src.parent.resolve() != data_dir.resolve():
1683
+ dst = data_dir / file_name
1684
+ if not dst.exists() or dst.stat().st_size != src.stat().st_size:
1685
+ try:
1686
+ shutil.copy2(src, dst)
1687
+ except Exception as exc:
1688
+ logger.warning("Copy failed for %s: %s", src, exc)
1689
+ continue
1690
+ stored_path = f"app/backend/data/{file_name}"
1691
+ else:
1692
+ stored_path = str(src)
1693
+
1694
+ existing_files[file_name] = {
1695
+ 'file_name': file_name,
1696
+ 'prompt': prompt,
1697
+ 'path': stored_path,
1698
+ }
1699
+ committed += 1
1700
+
1701
+ final_metadata = list(existing_files.values())
1702
+ with open(json_path, 'w', encoding='utf-8') as f:
1703
+ json.dump(final_metadata, f, indent=2)
1704
+
1705
+ try:
1706
+ config.update_dataset_config()
1707
+ except Exception as exc:
1708
+ logger.warning("Failed to refresh dataset-config.json: %s", exc)
1709
+
1710
+ return jsonify({
1711
+ 'message': f'Committed {committed} annotations.',
1712
+ 'committed': committed,
1713
+ 'metadata_json': str(json_path),
1714
+ })
1715
+
1716
+
1717
+ @app.route('/api/bulk-annotate/download-clap', methods=['POST'])
1718
+ def bulk_annotate_download_clap():
1719
+ from app.backend.data.auto_annotator import download_clap_checkpoint
1720
+
1721
+ with _annotate_job_lock:
1722
+ if _clap_download_job['state'] == 'running':
1723
+ return jsonify({'error': 'CLAP download already in progress.'}), 409
1724
+ _clap_download_job.update({'state': 'running', 'message': 'Starting download…', 'error': None})
1725
+
1726
+ def runner():
1727
+ try:
1728
+ target = download_clap_checkpoint(
1729
+ get_config().get_path('models_pretrained'),
1730
+ progress_cb=lambda m: _clap_download_job.update({'message': m}),
1731
+ )
1732
+ _clap_download_job.update({'state': 'done', 'message': f'Downloaded to {target}'})
1733
+ except Exception as exc:
1734
+ logger.exception("CLAP download failed")
1735
+ _clap_download_job.update({'state': 'error', 'error': str(exc)})
1736
+
1737
+ threading.Thread(target=runner, daemon=True).start()
1738
+ return jsonify({'message': 'CLAP download started'})
1739
+
1740
+
1741
+ @app.route('/api/bulk-annotate/unload-clap', methods=['POST'])
1742
+ def bulk_annotate_unload_clap():
1743
+ from app.backend.data.auto_annotator import unload_clap
1744
+ unload_clap()
1745
+ return jsonify({'message': 'CLAP unloaded from memory.'})
1746
+
1747
+
1748
  @app.route('/shutdown', methods=['POST'])
1749
  def shutdown():
1750
  """Shutdown the Flask server gracefully"""
 
1761
 
1762
 
1763
  if __name__ == '__main__':
1764
+ # 0.0.0.0: reachable at this machine's LAN/Tailscale IPs (e.g. http://100.122.31.32:5001).
1765
+ host = os.environ.get('FLASK_HOST', '0.0.0.0')
1766
+ port = int(os.environ.get('FLASK_PORT', '5001'))
1767
+ app.run(debug=True, host=host, port=port)
app/backend/data/auto_annotator.py ADDED
@@ -0,0 +1,406 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Automatic audio annotation for bulk dataset creation.
2
+
3
+ Two tiers:
4
+ - basic: librosa-only DSP (tempo, key). No downloads. CPU. ~instant per file.
5
+ - rich: basic + LAION-CLAP zero-shot tagging (genre, mood, instrument).
6
+ Lazy-loaded; downloads ~2.35 GB checkpoint on first use.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ import logging
13
+ import os
14
+ import threading
15
+ from pathlib import Path
16
+ from typing import Any, Callable, Dict, List, Optional
17
+
18
+ logger = logging.getLogger(__name__)
19
+
20
+ AUDIO_EXTENSIONS = (".wav", ".mp3", ".flac", ".m4a", ".ogg", ".aac")
21
+
22
+ CLAP_CKPT_FILENAME = "music_audioset_epoch_15_esc_90.14.pt"
23
+ CLAP_REPO = "lukewys/laion_clap"
24
+
25
+ KEY_NAMES_SHARP = ["C", "C#", "D", "D#", "E", "F", "F#", "G", "G#", "A", "A#", "B"]
26
+ KEY_NAMES_FLAT = ["C", "Db", "D", "Eb", "E", "F", "Gb", "G", "Ab", "A", "Bb", "B"]
27
+
28
+ # Krumhansl-Schmuckler key profiles.
29
+ KRUMHANSL_MAJOR = [6.35, 2.23, 3.48, 2.33, 4.38, 4.09, 2.52, 5.19, 2.39, 3.66, 2.29, 2.88]
30
+ KRUMHANSL_MINOR = [6.33, 2.68, 3.52, 5.38, 2.60, 3.53, 2.54, 4.75, 3.98, 2.69, 3.34, 3.17]
31
+
32
+
33
+ def _iter_audio_files(folder: Path) -> List[Path]:
34
+ results: List[Path] = []
35
+ for root, _, files in os.walk(folder):
36
+ for name in files:
37
+ if name.startswith("."):
38
+ continue
39
+ if name.lower().endswith(AUDIO_EXTENSIONS):
40
+ results.append(Path(root) / name)
41
+ results.sort()
42
+ return results
43
+
44
+
45
+ def _estimate_tempo(y, sr) -> Optional[int]:
46
+ import librosa
47
+ try:
48
+ tempo, _ = librosa.beat.beat_track(y=y, sr=sr)
49
+ bpm = float(tempo if hasattr(tempo, "__float__") else tempo[0])
50
+ if bpm <= 0:
51
+ return None
52
+ return int(round(bpm))
53
+ except Exception as exc:
54
+ logger.debug("tempo estimation failed: %s", exc)
55
+ return None
56
+
57
+
58
+ def _estimate_brightness(y, sr) -> Optional[str]:
59
+ import librosa
60
+ try:
61
+ centroid = float(librosa.feature.spectral_centroid(y=y, sr=sr).mean())
62
+ except Exception as exc:
63
+ logger.debug("centroid estimation failed: %s", exc)
64
+ return None
65
+ if centroid <= 0:
66
+ return None
67
+ if centroid < 1500:
68
+ return "dark"
69
+ if centroid > 3500:
70
+ return "bright"
71
+ return None
72
+
73
+
74
+ def _estimate_character(y, sr) -> Optional[str]:
75
+ import librosa
76
+ import numpy as np
77
+ try:
78
+ harm, perc = librosa.effects.hpss(y)
79
+ eh = float(np.mean(harm ** 2))
80
+ ep = float(np.mean(perc ** 2))
81
+ except Exception as exc:
82
+ logger.debug("HPSS failed: %s", exc)
83
+ return None
84
+ total = eh + ep
85
+ if total <= 0:
86
+ return None
87
+ perc_ratio = ep / total
88
+ if perc_ratio > 0.65:
89
+ return "percussion-driven"
90
+ if perc_ratio < 0.20:
91
+ return "melodic"
92
+ return None
93
+
94
+
95
+ def _estimate_key(y, sr) -> Optional[str]:
96
+ import librosa
97
+ import numpy as np
98
+ try:
99
+ chroma = librosa.feature.chroma_cqt(y=y, sr=sr)
100
+ chroma_mean = chroma.mean(axis=1)
101
+ if chroma_mean.sum() <= 0:
102
+ return None
103
+ chroma_mean = chroma_mean / chroma_mean.sum()
104
+
105
+ major = np.asarray(KRUMHANSL_MAJOR)
106
+ minor = np.asarray(KRUMHANSL_MINOR)
107
+
108
+ best_score = -1.0
109
+ best_key = None
110
+ for i in range(12):
111
+ maj_score = float(np.corrcoef(chroma_mean, np.roll(major, i))[0, 1])
112
+ min_score = float(np.corrcoef(chroma_mean, np.roll(minor, i))[0, 1])
113
+ if maj_score > best_score:
114
+ best_score = maj_score
115
+ best_key = f"{KEY_NAMES_SHARP[i]} major"
116
+ if min_score > best_score:
117
+ best_score = min_score
118
+ best_key = f"{KEY_NAMES_SHARP[i]} minor"
119
+ return best_key
120
+ except Exception as exc:
121
+ logger.debug("key estimation failed: %s", exc)
122
+ return None
123
+
124
+
125
+ def _compose_prompt(parts: Dict[str, Any]) -> str:
126
+ genre = parts.get("genre")
127
+ mood = parts.get("mood")
128
+ instruments = parts.get("instruments") or []
129
+ bpm = parts.get("bpm")
130
+ key = parts.get("key")
131
+ brightness = parts.get("brightness")
132
+ character = parts.get("character")
133
+
134
+ head_bits: List[str] = []
135
+ if mood:
136
+ head_bits.append(str(mood))
137
+ if genre:
138
+ head_bits.append(f"{genre} track")
139
+ elif head_bits:
140
+ head_bits[-1] = f"{head_bits[-1]} track"
141
+ opening = " ".join(head_bits)
142
+
143
+ descriptors = [d for d in (brightness, character) if d]
144
+
145
+ fragments: List[str] = []
146
+ if opening:
147
+ fragments.append(opening)
148
+ if descriptors:
149
+ fragments.append(", ".join(descriptors))
150
+ if bpm:
151
+ fragments.append(f"{bpm} BPM")
152
+ if key:
153
+ fragments.append(f"in {key}")
154
+ if instruments:
155
+ fragments.append("with " + ", ".join(instruments))
156
+
157
+ out = ", ".join(fragments)
158
+ return out[:1].upper() + out[1:] if out else ""
159
+
160
+
161
+ class _ClapTagger:
162
+ """Lazy holder for a LAION-CLAP model used for zero-shot tagging."""
163
+
164
+ def __init__(self, ckpt_path: Path):
165
+ self.ckpt_path = ckpt_path
166
+ self._model = None
167
+ self._lock = threading.Lock()
168
+ self._label_embeds: Dict[str, Any] = {}
169
+
170
+ def ensure_loaded(self):
171
+ if self._model is not None:
172
+ return
173
+ with self._lock:
174
+ if self._model is not None:
175
+ return
176
+ if not self.ckpt_path.exists():
177
+ raise FileNotFoundError(
178
+ f"CLAP checkpoint not found at {self.ckpt_path}. "
179
+ "Download it first via /api/bulk-annotate/download-clap."
180
+ )
181
+ import laion_clap
182
+ import torch
183
+ logging.getLogger("transformers").setLevel(logging.ERROR)
184
+
185
+ device = "cuda" if torch.cuda.is_available() else "cpu"
186
+ model = laion_clap.CLAP_Module(enable_fusion=False, amodel="HTSAT-base", device=device)
187
+
188
+ # torch >= 2.6 flipped torch.load(weights_only=True) and newer
189
+ # transformers dropped the roberta position_ids buffer, so
190
+ # laion_clap's own load_ckpt errors twice: unpickling, then strict
191
+ # state_dict mismatch. Replicate its logic safely here.
192
+ from laion_clap.clap_module.factory import load_state_dict as clap_load_state_dict
193
+
194
+ orig_load = torch.load
195
+ def _trusted_load(*args, **kwargs):
196
+ kwargs.setdefault("weights_only", False)
197
+ return orig_load(*args, **kwargs)
198
+ torch.load = _trusted_load
199
+ try:
200
+ state = clap_load_state_dict(str(self.ckpt_path), skip_params=True)
201
+ finally:
202
+ torch.load = orig_load
203
+ missing, unexpected = model.model.load_state_dict(state, strict=False)
204
+ if unexpected:
205
+ logger.debug("CLAP unexpected keys ignored: %s", unexpected[:5])
206
+ if missing:
207
+ logger.debug("CLAP missing keys: %s", missing[:5])
208
+ self._model = model
209
+ self._device = device
210
+ logger.info("CLAP loaded on %s from %s", device, self.ckpt_path)
211
+
212
+ def _embed_labels(self, group: str, prompts: List[str]):
213
+ import torch
214
+ key = f"{group}:{'|'.join(prompts)}"
215
+ if key in self._label_embeds:
216
+ return self._label_embeds[key]
217
+ with torch.no_grad():
218
+ embed = self._model.get_text_embedding(prompts, use_tensor=True)
219
+ embed = embed / embed.norm(dim=-1, keepdim=True).clamp_min(1e-8)
220
+ self._label_embeds[key] = embed
221
+ return embed
222
+
223
+ def tag(self, audio_path: Path, label_sets: Dict[str, List[str]], top_k_instruments: int = 2) -> Dict[str, Any]:
224
+ self.ensure_loaded()
225
+ import torch
226
+
227
+ with torch.no_grad():
228
+ audio_embed = self._model.get_audio_embedding_from_filelist(
229
+ x=[str(audio_path)], use_tensor=True
230
+ )
231
+ audio_embed = audio_embed / audio_embed.norm(dim=-1, keepdim=True).clamp_min(1e-8)
232
+
233
+ out: Dict[str, Any] = {}
234
+ for group in ("genre", "mood"):
235
+ labels = label_sets.get(group) or []
236
+ if not labels:
237
+ continue
238
+ prompts = [f"a {lab} music track" if group == "genre" else f"a {lab} sounding music track" for lab in labels]
239
+ text_embed = self._embed_labels(group, prompts)
240
+ sims = (audio_embed @ text_embed.T).squeeze(0)
241
+ top = int(sims.argmax().item())
242
+ out[group] = labels[top]
243
+
244
+ instruments = label_sets.get("instruments") or []
245
+ if instruments:
246
+ prompts = [f"music featuring {lab}" for lab in instruments]
247
+ text_embed = self._embed_labels("instruments", prompts)
248
+ sims = (audio_embed @ text_embed.T).squeeze(0)
249
+ k = min(top_k_instruments, len(instruments))
250
+ top_idx = torch.topk(sims, k=k).indices.tolist()
251
+ out["instruments"] = [instruments[i] for i in top_idx]
252
+ return out
253
+
254
+
255
+ _clap_tagger_singleton: Optional[_ClapTagger] = None
256
+ _clap_tagger_lock = threading.Lock()
257
+
258
+
259
+ def get_clap_tagger(clap_ckpt_path: Path) -> _ClapTagger:
260
+ global _clap_tagger_singleton
261
+ with _clap_tagger_lock:
262
+ if _clap_tagger_singleton is None or _clap_tagger_singleton.ckpt_path != clap_ckpt_path:
263
+ _clap_tagger_singleton = _ClapTagger(clap_ckpt_path)
264
+ return _clap_tagger_singleton
265
+
266
+
267
+ def clap_checkpoint_path(models_pretrained_dir: Path) -> Path:
268
+ return models_pretrained_dir / "clap" / CLAP_CKPT_FILENAME
269
+
270
+
271
+ def clap_checkpoint_available(models_pretrained_dir: Path) -> bool:
272
+ return clap_checkpoint_path(models_pretrained_dir).exists()
273
+
274
+
275
+ def download_clap_checkpoint(
276
+ models_pretrained_dir: Path,
277
+ progress_cb: Optional[Callable[[str], None]] = None,
278
+ ) -> Path:
279
+ target = clap_checkpoint_path(models_pretrained_dir)
280
+ target.parent.mkdir(parents=True, exist_ok=True)
281
+ if target.exists():
282
+ return target
283
+
284
+ from huggingface_hub import hf_hub_download
285
+
286
+ if progress_cb:
287
+ progress_cb("Downloading CLAP checkpoint (~630 MB)…")
288
+
289
+ downloaded = hf_hub_download(
290
+ repo_id=CLAP_REPO,
291
+ filename=CLAP_CKPT_FILENAME,
292
+ local_dir=str(target.parent),
293
+ )
294
+ downloaded_path = Path(downloaded)
295
+ if downloaded_path != target:
296
+ try:
297
+ downloaded_path.replace(target)
298
+ except OSError:
299
+ import shutil
300
+ shutil.copy2(downloaded_path, target)
301
+ return target
302
+
303
+
304
+ def load_label_sets(label_sets_path: Optional[Path]) -> Dict[str, List[str]]:
305
+ if not label_sets_path or not label_sets_path.exists():
306
+ return {"genre": [], "mood": [], "instruments": []}
307
+ with open(label_sets_path, "r", encoding="utf-8") as f:
308
+ data = json.load(f)
309
+ return {
310
+ "genre": list(data.get("genre") or []),
311
+ "mood": list(data.get("mood") or []),
312
+ "instruments": list(data.get("instruments") or []),
313
+ }
314
+
315
+
316
+ def annotate_file(
317
+ audio_path: Path,
318
+ tier: str,
319
+ clap_tagger: Optional[_ClapTagger],
320
+ label_sets: Dict[str, List[str]],
321
+ sr: int = 22050,
322
+ max_seconds: float = 60.0,
323
+ ) -> Dict[str, Any]:
324
+ import librosa
325
+
326
+ parts: Dict[str, Any] = {}
327
+ try:
328
+ y, loaded_sr = librosa.load(str(audio_path), sr=sr, mono=True, duration=max_seconds)
329
+ except Exception as exc:
330
+ logger.warning("librosa failed to load %s: %s", audio_path.name, exc)
331
+ return {
332
+ "file_name": audio_path.name,
333
+ "prompt": "",
334
+ "path": str(audio_path),
335
+ "error": f"load failed: {exc}",
336
+ }
337
+
338
+ parts["bpm"] = _estimate_tempo(y, loaded_sr)
339
+ parts["key"] = _estimate_key(y, loaded_sr)
340
+ parts["brightness"] = _estimate_brightness(y, loaded_sr)
341
+ parts["character"] = _estimate_character(y, loaded_sr)
342
+
343
+ if tier == "rich" and clap_tagger is not None:
344
+ try:
345
+ tags = clap_tagger.tag(audio_path, label_sets)
346
+ parts.update(tags)
347
+ except Exception as exc:
348
+ logger.warning("CLAP tagging failed for %s: %s", audio_path.name, exc)
349
+
350
+ prompt = _compose_prompt(parts)
351
+ return {
352
+ "file_name": audio_path.name,
353
+ "prompt": prompt,
354
+ "path": str(audio_path),
355
+ "attributes": parts,
356
+ }
357
+
358
+
359
+ def annotate_folder(
360
+ folder: Path,
361
+ tier: str,
362
+ label_sets: Dict[str, List[str]],
363
+ clap_ckpt_path: Optional[Path] = None,
364
+ progress_cb: Optional[Callable[[int, int, str], None]] = None,
365
+ ) -> List[Dict[str, Any]]:
366
+ folder = Path(folder)
367
+ if not folder.exists() or not folder.is_dir():
368
+ raise ValueError(f"Folder not found: {folder}")
369
+
370
+ files = _iter_audio_files(folder)
371
+ if not files:
372
+ raise ValueError(f"No audio files found in {folder}")
373
+
374
+ clap_tagger: Optional[_ClapTagger] = None
375
+ if tier == "rich":
376
+ if not clap_ckpt_path or not Path(clap_ckpt_path).exists():
377
+ raise FileNotFoundError(
378
+ "Rich tier requires the CLAP checkpoint; download it first."
379
+ )
380
+ clap_tagger = get_clap_tagger(Path(clap_ckpt_path))
381
+ clap_tagger.ensure_loaded()
382
+
383
+ results: List[Dict[str, Any]] = []
384
+ total = len(files)
385
+ for i, audio_path in enumerate(files, start=1):
386
+ if progress_cb:
387
+ progress_cb(i, total, audio_path.name)
388
+ entry = annotate_file(audio_path, tier, clap_tagger, label_sets)
389
+ results.append(entry)
390
+ return results
391
+
392
+
393
+ def unload_clap():
394
+ """Free CLAP weights from VRAM. Call before training starts."""
395
+ global _clap_tagger_singleton
396
+ with _clap_tagger_lock:
397
+ if _clap_tagger_singleton is not None:
398
+ _clap_tagger_singleton._model = None
399
+ _clap_tagger_singleton._label_embeds = {}
400
+ _clap_tagger_singleton = None
401
+ try:
402
+ import torch
403
+ if torch.cuda.is_available():
404
+ torch.cuda.empty_cache()
405
+ except Exception:
406
+ pass
app/core/config.py CHANGED
@@ -8,14 +8,18 @@ class ProjectConfig:
8
 
9
  def __init__(self, project_root: Optional[Path] = None) -> None:
10
  if getattr(sys, 'frozen', False):
 
11
  self.frozen = True
 
12
  self.project_root = Path(sys._MEIPASS)
13
 
 
14
  if sys.platform == "win32":
15
  self.user_data_dir = Path(os.environ["APPDATA"]) / "FragmentaDesktop"
16
  elif sys.platform == "darwin":
17
  self.user_data_dir = Path.home() / "Library" / "Application Support" / "FragmentaDesktop"
18
  else:
 
19
  self.user_data_dir = Path.home() / ".local" / "share" / "FragmentaDesktop"
20
 
21
  self.user_data_dir.mkdir(parents=True, exist_ok=True)
@@ -41,6 +45,7 @@ class ProjectConfig:
41
  self.user_data_dir = self.project_root
42
 
43
  self.paths: Dict[str, Path] = {
 
44
  "models": self.user_data_dir / "models",
45
  "models_config": self.user_data_dir / "models" / "config",
46
  "models_pretrained": self.user_data_dir / "models" / "pretrained",
@@ -134,7 +139,7 @@ class ProjectConfig:
134
  "datasets": [
135
  {
136
  "id": "fine_tune_data",
137
- "path": str(self.paths["data"]),
138
  "custom_metadata_module": "custom_metadata"
139
  }
140
  ],
 
8
 
9
  def __init__(self, project_root: Optional[Path] = None) -> None:
10
  if getattr(sys, 'frozen', False):
11
+ # Running in PyInstaller bundle
12
  self.frozen = True
13
+ # sys._MEIPASS is where PyInstaller unpacks the bundle
14
  self.project_root = Path(sys._MEIPASS)
15
 
16
+ # For writable data, use a user directory
17
  if sys.platform == "win32":
18
  self.user_data_dir = Path(os.environ["APPDATA"]) / "FragmentaDesktop"
19
  elif sys.platform == "darwin":
20
  self.user_data_dir = Path.home() / "Library" / "Application Support" / "FragmentaDesktop"
21
  else:
22
+ # Linux/Unix
23
  self.user_data_dir = Path.home() / ".local" / "share" / "FragmentaDesktop"
24
 
25
  self.user_data_dir.mkdir(parents=True, exist_ok=True)
 
45
  self.user_data_dir = self.project_root
46
 
47
  self.paths: Dict[str, Path] = {
48
+ # Writable paths - go to user_data_dir in frozen mode
49
  "models": self.user_data_dir / "models",
50
  "models_config": self.user_data_dir / "models" / "config",
51
  "models_pretrained": self.user_data_dir / "models" / "pretrained",
 
139
  "datasets": [
140
  {
141
  "id": "fine_tune_data",
142
+ "path": str(self.paths["data"].relative_to(self.project_root)),
143
  "custom_metadata_module": "custom_metadata"
144
  }
145
  ],
app/core/generation/audio_generator.py CHANGED
@@ -1,15 +1,23 @@
1
  import torch
2
- import torchaudio
3
  import numpy as np
4
  from pathlib import Path
5
  from typing import Dict, Any, Optional, List, Tuple
6
  import logging
7
  import sys
8
  import time
 
9
 
10
  sys.path.append(
11
  str(Path(__file__).parent.parent.parent.parent / "stable-audio-tools"))
12
 
 
 
 
 
 
 
 
13
  from stable_audio_tools.models.utils import load_ckpt_state_dict
14
  from stable_audio_tools.inference.generation import generate_diffusion_cond
15
  from stable_audio_tools.models import create_model_from_config
@@ -68,10 +76,7 @@ class AudioGenerator:
68
  self.model.eval()
69
  self.model.requires_grad_(False)
70
  if self.device.startswith("cuda"):
71
- try:
72
- self.model = torch.compile(self.model, mode="reduce-overhead")
73
- except Exception as compile_err:
74
- logger.warning(f"torch.compile failed ({compile_err}), using eager mode")
75
 
76
  logger.info("Local base model loaded successfully")
77
  return True
@@ -129,10 +134,7 @@ class AudioGenerator:
129
  self.model.requires_grad_(False)
130
 
131
  if self.device.startswith("cuda"):
132
- try:
133
- self.model = torch.compile(self.model, mode="reduce-overhead")
134
- except Exception as compile_err:
135
- logger.warning(f"torch.compile failed ({compile_err}), using eager mode")
136
 
137
  print(f"AUDIO GENERATOR: Unwrapped model loaded successfully")
138
  return True
@@ -174,23 +176,7 @@ class AudioGenerator:
174
  print(f" - Prompt: '{prompt}'")
175
  print(f" - Duration: {duration}s")
176
 
177
- # Determine what model we need to load
178
- needs_load = False
179
- if self.model is None:
180
- print(f"AUDIO GENERATOR: No model loaded yet")
181
- needs_load = True
182
- elif unwrapped_model_path:
183
- # Unwrapped model requested — reload if different
184
- if self.current_model_path != unwrapped_model_path:
185
- print(f"AUDIO GENERATOR: Switching to different unwrapped model")
186
- needs_load = True
187
- elif model_path is not None:
188
- # Base/fine-tuned model requested — reload only if different
189
- if self.current_model_path != str(model_path):
190
- print(f"AUDIO GENERATOR: Switching to different model (current={self.current_model_path}, requested={model_path})")
191
- needs_load = True
192
-
193
- if needs_load:
194
  print(f"AUDIO GENERATOR: Loading new model")
195
 
196
  if unwrapped_model_path:
@@ -279,19 +265,34 @@ class AudioGenerator:
279
  device = next(self.model.parameters()).device
280
  print(f"Using device: {device}")
281
 
282
- audio = generate_diffusion_cond(
283
- model=self.model,
284
- steps=steps,
285
- cfg_scale=cfg_scale,
286
- conditioning=conditioning,
287
- batch_size=1,
288
- sample_size=requested_sample_size,
289
- seed=seed,
290
- device=str(device),
291
- sigma_min=0.03,
292
- sigma_max=1000,
293
- sampler_type="dpmpp-3m-sde"
294
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
295
 
296
  print(f"Generation complete, audio shape: {audio.shape}")
297
 
@@ -358,7 +359,8 @@ class AudioGenerator:
358
 
359
  def save_audio(self, audio: torch.Tensor, output_path: Path, sample_rate: int):
360
  output_path.parent.mkdir(exist_ok=True, parents=True)
361
- torchaudio.save(str(output_path), audio, sample_rate)
 
362
 
363
  def get_model_info(self) -> Dict[str, Any]:
364
  if self.model is None:
 
1
  import torch
2
+ import soundfile as sf
3
  import numpy as np
4
  from pathlib import Path
5
  from typing import Dict, Any, Optional, List, Tuple
6
  import logging
7
  import sys
8
  import time
9
+ import warnings
10
 
11
  sys.path.append(
12
  str(Path(__file__).parent.parent.parent.parent / "stable-audio-tools"))
13
 
14
+ # Third-party package noise (clip/pkg_resources) is non-actionable for runtime.
15
+ warnings.filterwarnings(
16
+ "ignore",
17
+ message=r"pkg_resources is deprecated as an API.*",
18
+ category=UserWarning,
19
+ )
20
+
21
  from stable_audio_tools.models.utils import load_ckpt_state_dict
22
  from stable_audio_tools.inference.generation import generate_diffusion_cond
23
  from stable_audio_tools.models import create_model_from_config
 
76
  self.model.eval()
77
  self.model.requires_grad_(False)
78
  if self.device.startswith("cuda"):
79
+ self.model = torch.compile(self.model, mode="reduce-overhead")
 
 
 
80
 
81
  logger.info("Local base model loaded successfully")
82
  return True
 
134
  self.model.requires_grad_(False)
135
 
136
  if self.device.startswith("cuda"):
137
+ self.model = torch.compile(self.model, mode="reduce-overhead")
 
 
 
138
 
139
  print(f"AUDIO GENERATOR: Unwrapped model loaded successfully")
140
  return True
 
176
  print(f" - Prompt: '{prompt}'")
177
  print(f" - Duration: {duration}s")
178
 
179
+ if self.model is None or model_path is not None or unwrapped_model_path is not None:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
180
  print(f"AUDIO GENERATOR: Loading new model")
181
 
182
  if unwrapped_model_path:
 
265
  device = next(self.model.parameters()).device
266
  print(f"Using device: {device}")
267
 
268
+ with warnings.catch_warnings():
269
+ # Known torchsde float-boundary chatter from dpmpp-3m-sde.
270
+ warnings.filterwarnings(
271
+ "ignore",
272
+ message=r"Should have tb<=t1 but got tb=.*",
273
+ category=UserWarning,
274
+ module=r"torchsde\._brownian\.brownian_interval",
275
+ )
276
+ warnings.filterwarnings(
277
+ "ignore",
278
+ message=r"Should have ta>=t0 but got ta=.*",
279
+ category=UserWarning,
280
+ module=r"torchsde\._brownian\.brownian_interval",
281
+ )
282
+
283
+ audio = generate_diffusion_cond(
284
+ model=self.model,
285
+ steps=steps,
286
+ cfg_scale=cfg_scale,
287
+ conditioning=conditioning,
288
+ batch_size=1,
289
+ sample_size=requested_sample_size,
290
+ seed=seed,
291
+ device=str(device),
292
+ sigma_min=0.03,
293
+ sigma_max=1000,
294
+ sampler_type="dpmpp-3m-sde"
295
+ )
296
 
297
  print(f"Generation complete, audio shape: {audio.shape}")
298
 
 
359
 
360
  def save_audio(self, audio: torch.Tensor, output_path: Path, sample_rate: int):
361
  output_path.parent.mkdir(exist_ok=True, parents=True)
362
+ audio_np = audio.detach().cpu().transpose(0, 1).numpy()
363
+ sf.write(str(output_path), audio_np, sample_rate, subtype="PCM_16")
364
 
365
  def get_model_info(self) -> Dict[str, Any]:
366
  if self.model is None:
app/core/model_manager.py CHANGED
@@ -203,50 +203,17 @@ class ModelManager:
203
  print(f"Not authenticated with Hugging Face: {auth_error}")
204
  if progress_callback:
205
  progress_callback(0, "Authentication required...")
206
-
207
- is_docker = os.environ.get('FRAGMENTA_DOCKER', '').strip() == '1'
208
- if is_docker:
209
- print("Docker mode: HF authentication required. "
210
- "Set your token via Model Setup in the browser UI, "
211
- "or pass -e HF_TOKEN=hf_xxx to docker run.")
212
- if progress_callback:
213
- progress_callback(0, "HF authentication required — use Model Setup to set your token")
214
- return False
215
-
216
- try:
217
- from app.core.hf_auth_dialog import show_hf_auth_dialog
218
- success = show_hf_auth_dialog()
219
-
220
- if not success:
221
- print("Authentication dialog was cancelled")
222
- if progress_callback:
223
- progress_callback(0, "Authentication cancelled")
224
- return False
225
-
226
- try:
227
- user = api.whoami()
228
- print(f"Now authenticated as: {user}")
229
- if progress_callback:
230
- progress_callback(
231
- 10, "Authentication successful...")
232
- except Exception as retry_error:
233
- print(f"Still not authenticated: {retry_error}")
234
- if progress_callback:
235
- progress_callback(0, "Authentication failed")
236
- return False
237
-
238
- except ImportError:
239
- print("To download models, you need to:")
240
- print(
241
- "1. Visit https://huggingface.co/stabilityai/stable-audio-open-small")
242
- print("2. Accept the terms and conditions")
243
- print("3. Log in to your Hugging Face account")
244
- print(
245
- "4. Get your access token from https://huggingface.co/settings/tokens")
246
- print("5. Run: huggingface-cli login")
247
- if progress_callback:
248
- progress_callback(0, "Manual authentication required")
249
- return False
250
 
251
  if progress_callback:
252
  progress_callback(20, "Starting file download...")
@@ -257,6 +224,7 @@ class ModelManager:
257
  from tqdm import tqdm
258
  import sys
259
 
 
260
  class TqdmToCallback:
261
  def __init__(self, callback, file_index, total_files):
262
  self.callback = callback
@@ -268,6 +236,7 @@ class ModelManager:
268
  """Returns a callback function for tqdm"""
269
  def inner(bytes_amount=1):
270
  if t.total:
 
271
  file_progress = (t.n / t.total)
272
  overall_progress = (self.file_index + file_progress) / self.total_files
273
  percent = 20 + int(overall_progress * 70)
@@ -304,8 +273,10 @@ class ModelManager:
304
  else:
305
  final_filename = f"{model_id}-{file_pattern}"
306
 
 
307
  tqdm_callback = TqdmToCallback(progress_callback, i, total_files)
308
 
 
309
  original_tqdm_init = tqdm.__init__
310
 
311
  def patched_tqdm_init(self, *args, **kwargs):
@@ -336,6 +307,7 @@ class ModelManager:
336
  resume_download=True
337
  )
338
  finally:
 
339
  tqdm.__init__ = original_tqdm_init
340
 
341
  downloaded_path = Path(downloaded_file)
 
203
  print(f"Not authenticated with Hugging Face: {auth_error}")
204
  if progress_callback:
205
  progress_callback(0, "Authentication required...")
206
+ print("To download models, you need to:")
207
+ print(
208
+ "1. Visit https://huggingface.co/stabilityai/stable-audio-open-small")
209
+ print("2. Accept the terms and conditions")
210
+ print("3. Log in to your Hugging Face account")
211
+ print(
212
+ "4. Get your access token from https://huggingface.co/settings/tokens")
213
+ print("5. Use the in-app Hugging Face login dialog")
214
+ if progress_callback:
215
+ progress_callback(0, "Please authenticate in the app first")
216
+ return False
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
217
 
218
  if progress_callback:
219
  progress_callback(20, "Starting file download...")
 
224
  from tqdm import tqdm
225
  import sys
226
 
227
+ # Redirect tqdm to capture progress
228
  class TqdmToCallback:
229
  def __init__(self, callback, file_index, total_files):
230
  self.callback = callback
 
236
  """Returns a callback function for tqdm"""
237
  def inner(bytes_amount=1):
238
  if t.total:
239
+ # Calculate progress: 20-90% range for all files
240
  file_progress = (t.n / t.total)
241
  overall_progress = (self.file_index + file_progress) / self.total_files
242
  percent = 20 + int(overall_progress * 70)
 
273
  else:
274
  final_filename = f"{model_id}-{file_pattern}"
275
 
276
+ # Use custom tqdm callback to intercept progress
277
  tqdm_callback = TqdmToCallback(progress_callback, i, total_files)
278
 
279
+ # Monkey-patch tqdm for this download
280
  original_tqdm_init = tqdm.__init__
281
 
282
  def patched_tqdm_init(self, *args, **kwargs):
 
307
  resume_download=True
308
  )
309
  finally:
310
+ # Restore original tqdm
311
  tqdm.__init__ = original_tqdm_init
312
 
313
  downloaded_path = Path(downloaded_file)
app/core/training/lora_trainer.py CHANGED
@@ -10,6 +10,9 @@ import torch.mps
10
 
11
  from app.core.config import get_config
12
 
 
 
 
13
  def get_base_model_configs():
14
  config = get_config()
15
  return config.model_configs
@@ -24,7 +27,7 @@ class LoRATrainer:
24
  "is_training": False,
25
  "progress": 0,
26
  "current_epoch": 0,
27
- "total_epochs": config.get("epochs", 10),
28
  "loss": None,
29
  "loss_history": [],
30
  "error": None,
@@ -37,6 +40,26 @@ class LoRATrainer:
37
  "checkpoints_saved": 0,
38
  }
39
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
40
  def _validate_dataset_before_training(self) -> Dict[str, Any]:
41
  try:
42
  from app.core.config import get_config
@@ -71,16 +94,16 @@ class LoRATrainer:
71
  }
72
 
73
  steps_per_epoch = max(1, file_count // batch_size)
74
- total_epochs = self.config.get("epochs", 10)
75
  total_steps = steps_per_epoch * total_epochs
76
- checkpoint_every = self.config.get("checkpointSteps", 100)
77
-
78
- checkpoint_warning = None
79
- if total_steps < checkpoint_every:
80
- checkpoint_warning = (
81
- f"WARNING: NOT ENOUGH DATA! "
82
- f"Training will complete in {total_steps} steps, but checkpoints are set to save every {checkpoint_every} steps. "
83
- f"Add more audio files or reduce checkpoint interval."
84
  )
85
 
86
  return {
@@ -90,7 +113,8 @@ class LoRATrainer:
90
  "steps_per_epoch": steps_per_epoch,
91
  "total_steps": total_steps,
92
  "checkpoint_every": checkpoint_every,
93
- "checkpoint_warning": checkpoint_warning
 
94
  }
95
 
96
  except Exception as e:
@@ -126,7 +150,7 @@ class LoRATrainer:
126
  print(f"RECEIVED TRAINING CONFIG:")
127
  print(f"- Model Name: {self.config.get('modelName', 'untitled')}")
128
  print(f"- Base Model: {base_model}")
129
- print(f"- Epochs: {self.config.get('epochs', 3)}")
130
  print(f"- Batch Size: {self.config.get('batchSize', 1)}")
131
  print(f"- Learning Rate: {self.config.get('learningRate', 1e-4)}")
132
  base_model_configs = get_base_model_configs()
@@ -261,6 +285,28 @@ class LoRATrainer:
261
  if not os.path.exists(venv_python):
262
  venv_python = sys.executable
263
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
264
  cmd = [
265
  venv_python, "train.py",
266
  "--pretrained-ckpt-path", pretrained_ckpt,
@@ -268,7 +314,7 @@ class LoRATrainer:
268
  "--dataset-config", dataset_config,
269
  "--name", model_name,
270
  "--save-dir", save_dir,
271
- "--checkpoint-every", str(self.config.get("checkpointSteps", 25)),
272
  "--batch-size", str(batch_size),
273
  ] + memory_flags
274
 
@@ -300,35 +346,6 @@ class LoRATrainer:
300
  print(f"- Available: {device_info['memory_gb']:.2f} GB")
301
  print(f"- Using: {device_info['memory_ratio']*100:.0f}%")
302
 
303
- validation_result = self._validate_dataset_before_training()
304
- if not validation_result["valid"]:
305
- print(f"TRAINING VALIDATION FAILED!")
306
- print(f"{validation_result['error']}")
307
- return {
308
- "success": False,
309
- "error": validation_result["error"],
310
- "validation_failed": True
311
- }
312
-
313
- if validation_result.get("checkpoint_warning"):
314
- print(f"\nCHECKPOINT WARNING:")
315
- print(f"{validation_result['checkpoint_warning']}")
316
- print(f"Training Stats:")
317
- print(f"- Audio files: {validation_result['file_count']}")
318
- print(f"- Batch size: {validation_result['batch_size']}")
319
- print(f"- Steps per epoch: {validation_result['steps_per_epoch']}")
320
- print(f"- Total epochs: {validation_result['total_steps'] // validation_result['steps_per_epoch']}")
321
- print(f"- Total steps: {validation_result['total_steps']}")
322
- print(f"- Checkpoint every: {validation_result['checkpoint_every']} steps")
323
- print(f"Recommended checkpoint interval: {max(10, validation_result['total_steps'] // 2)} steps")
324
-
325
- return {
326
- "success": False,
327
- "error": validation_result["checkpoint_warning"],
328
- "checkpoint_warning": True,
329
- "validation_stats": validation_result
330
- }
331
-
332
  print(f"\nSTARTING TRAINING PROCESS...")
333
  print("="*80)
334
 
@@ -352,7 +369,7 @@ class LoRATrainer:
352
  audio_files.extend(list(data_dir.glob(f"*{ext}")))
353
 
354
  num_files = len(audio_files)
355
- total_epochs = self.config.get("epochs", 10)
356
  steps_per_epoch = max(1, num_files // batch_size)
357
  total_steps = steps_per_epoch * total_epochs
358
 
@@ -362,7 +379,7 @@ class LoRATrainer:
362
  print(f"- Steps per epoch: {steps_per_epoch}")
363
  print(f"- Total epochs: {total_epochs}")
364
  print(f"- Total steps: {total_steps}")
365
- print(f"- Checkpoint every: {self.config.get('checkpointSteps', 25)} steps")
366
 
367
  self.training_status.update({
368
  "is_training": True,
@@ -378,6 +395,7 @@ class LoRATrainer:
378
  "total_steps_per_epoch": steps_per_epoch, # Frontend expects this name
379
  "total_steps": total_steps,
380
  "total_epochs": total_epochs,
 
381
  "error_messages": [],
382
  "stop_initiated": False
383
  })
@@ -617,7 +635,7 @@ class LoRATrainer:
617
 
618
  if training_actually_succeeded:
619
  completed_epoch = self.training_status.get("current_epoch", 0)
620
- target_epochs = self.training_status.get("total_epochs", self.config.get("epochs", 10))
621
  total_steps = self.training_status.get("total_steps", 0)
622
 
623
  print(f"\n{'='*80}")
@@ -700,7 +718,7 @@ class LoRATrainer:
700
 
701
  steps_per_epoch = self.training_status.get("steps_per_epoch", 1)
702
  total_steps = self.training_status.get("total_steps", 1)
703
- target_epochs = self.training_status.get("total_epochs", 10)
704
 
705
  current_epoch = (global_step - 1) // steps_per_epoch
706
  current_step = ((global_step - 1) % steps_per_epoch) + 1
@@ -715,7 +733,10 @@ class LoRATrainer:
715
 
716
  if global_step >= total_steps:
717
  if not self.training_status.get("stop_initiated"):
718
- checkpoint_interval = self.config.get("checkpointSteps", 100)
 
 
 
719
 
720
  print(f"\n{'='*80}")
721
  print(f"TARGET EPOCHS REACHED: {current_epoch + 1}/{target_epochs}")
 
10
 
11
  from app.core.config import get_config
12
 
13
+ DEFAULT_EPOCHS = 30
14
+ DEFAULT_CHECKPOINT_STEPS = 50
15
+
16
  def get_base_model_configs():
17
  config = get_config()
18
  return config.model_configs
 
27
  "is_training": False,
28
  "progress": 0,
29
  "current_epoch": 0,
30
+ "total_epochs": config.get("epochs", DEFAULT_EPOCHS),
31
  "loss": None,
32
  "loss_history": [],
33
  "error": None,
 
40
  "checkpoints_saved": 0,
41
  }
42
 
43
+ @staticmethod
44
+ def _resolve_checkpoint_interval(total_steps: int, requested_interval: int) -> int:
45
+ """Choose a checkpoint interval that lands exactly on the final step."""
46
+ requested = max(1, int(requested_interval))
47
+ if total_steps <= 0:
48
+ return requested
49
+
50
+ requested = min(requested, total_steps)
51
+ if total_steps % requested == 0:
52
+ return requested
53
+
54
+ # Prefer a nearby lower divisor to keep cadence similar to user input.
55
+ min_reasonable = max(10, requested // 3)
56
+ for candidate in range(requested - 1, min_reasonable - 1, -1):
57
+ if total_steps % candidate == 0:
58
+ return candidate
59
+
60
+ # Fall back to one checkpoint at the exact final step.
61
+ return total_steps
62
+
63
  def _validate_dataset_before_training(self) -> Dict[str, Any]:
64
  try:
65
  from app.core.config import get_config
 
94
  }
95
 
96
  steps_per_epoch = max(1, file_count // batch_size)
97
+ total_epochs = self.config.get("epochs", DEFAULT_EPOCHS)
98
  total_steps = steps_per_epoch * total_epochs
99
+ requested_checkpoint_every = self.config.get("checkpointSteps", DEFAULT_CHECKPOINT_STEPS)
100
+ checkpoint_every = self._resolve_checkpoint_interval(total_steps, requested_checkpoint_every)
101
+
102
+ checkpoint_adjustment = None
103
+ if checkpoint_every != requested_checkpoint_every:
104
+ checkpoint_adjustment = (
105
+ f"Adjusted checkpoint interval from {requested_checkpoint_every} to {checkpoint_every} "
106
+ f"so the final step ({total_steps}) is always checkpointed."
107
  )
108
 
109
  return {
 
113
  "steps_per_epoch": steps_per_epoch,
114
  "total_steps": total_steps,
115
  "checkpoint_every": checkpoint_every,
116
+ "requested_checkpoint_every": requested_checkpoint_every,
117
+ "checkpoint_adjustment": checkpoint_adjustment,
118
  }
119
 
120
  except Exception as e:
 
150
  print(f"RECEIVED TRAINING CONFIG:")
151
  print(f"- Model Name: {self.config.get('modelName', 'untitled')}")
152
  print(f"- Base Model: {base_model}")
153
+ print(f"- Epochs: {self.config.get('epochs', DEFAULT_EPOCHS)}")
154
  print(f"- Batch Size: {self.config.get('batchSize', 1)}")
155
  print(f"- Learning Rate: {self.config.get('learningRate', 1e-4)}")
156
  base_model_configs = get_base_model_configs()
 
285
  if not os.path.exists(venv_python):
286
  venv_python = sys.executable
287
 
288
+ validation_result = self._validate_dataset_before_training()
289
+ if not validation_result["valid"]:
290
+ print(f"TRAINING VALIDATION FAILED!")
291
+ print(f"{validation_result['error']}")
292
+ return {
293
+ "success": False,
294
+ "error": validation_result["error"],
295
+ "validation_failed": True
296
+ }
297
+
298
+ checkpoint_every = validation_result["checkpoint_every"]
299
+ if validation_result.get("checkpoint_adjustment"):
300
+ print(f"\nCHECKPOINT INTERVAL NOTICE:")
301
+ print(f"{validation_result['checkpoint_adjustment']}")
302
+ print(f"Training Stats:")
303
+ print(f"- Audio files: {validation_result['file_count']}")
304
+ print(f"- Batch size: {validation_result['batch_size']}")
305
+ print(f"- Steps per epoch: {validation_result['steps_per_epoch']}")
306
+ print(f"- Total steps: {validation_result['total_steps']}")
307
+ print(f"- Requested checkpoint every: {validation_result['requested_checkpoint_every']} steps")
308
+ print(f"- Effective checkpoint every: {checkpoint_every} steps")
309
+
310
  cmd = [
311
  venv_python, "train.py",
312
  "--pretrained-ckpt-path", pretrained_ckpt,
 
314
  "--dataset-config", dataset_config,
315
  "--name", model_name,
316
  "--save-dir", save_dir,
317
+ "--checkpoint-every", str(checkpoint_every),
318
  "--batch-size", str(batch_size),
319
  ] + memory_flags
320
 
 
346
  print(f"- Available: {device_info['memory_gb']:.2f} GB")
347
  print(f"- Using: {device_info['memory_ratio']*100:.0f}%")
348
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
349
  print(f"\nSTARTING TRAINING PROCESS...")
350
  print("="*80)
351
 
 
369
  audio_files.extend(list(data_dir.glob(f"*{ext}")))
370
 
371
  num_files = len(audio_files)
372
+ total_epochs = self.config.get("epochs", DEFAULT_EPOCHS)
373
  steps_per_epoch = max(1, num_files // batch_size)
374
  total_steps = steps_per_epoch * total_epochs
375
 
 
379
  print(f"- Steps per epoch: {steps_per_epoch}")
380
  print(f"- Total epochs: {total_epochs}")
381
  print(f"- Total steps: {total_steps}")
382
+ print(f"- Checkpoint every: {checkpoint_every} steps")
383
 
384
  self.training_status.update({
385
  "is_training": True,
 
395
  "total_steps_per_epoch": steps_per_epoch, # Frontend expects this name
396
  "total_steps": total_steps,
397
  "total_epochs": total_epochs,
398
+ "checkpoint_every": checkpoint_every,
399
  "error_messages": [],
400
  "stop_initiated": False
401
  })
 
635
 
636
  if training_actually_succeeded:
637
  completed_epoch = self.training_status.get("current_epoch", 0)
638
+ target_epochs = self.training_status.get("total_epochs", self.config.get("epochs", DEFAULT_EPOCHS))
639
  total_steps = self.training_status.get("total_steps", 0)
640
 
641
  print(f"\n{'='*80}")
 
718
 
719
  steps_per_epoch = self.training_status.get("steps_per_epoch", 1)
720
  total_steps = self.training_status.get("total_steps", 1)
721
+ target_epochs = self.training_status.get("total_epochs", DEFAULT_EPOCHS)
722
 
723
  current_epoch = (global_step - 1) // steps_per_epoch
724
  current_step = ((global_step - 1) % steps_per_epoch) + 1
 
733
 
734
  if global_step >= total_steps:
735
  if not self.training_status.get("stop_initiated"):
736
+ checkpoint_interval = self.training_status.get(
737
+ "checkpoint_every",
738
+ self.config.get("checkpointSteps", DEFAULT_CHECKPOINT_STEPS)
739
+ )
740
 
741
  print(f"\n{'='*80}")
742
  print(f"TARGET EPOCHS REACHED: {current_epoch + 1}/{target_epochs}")
app/frontend/build/BitcountSingle-VariableFont_CRSV,ELSH,ELXP,slnt,wght.ttf ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c4d16dfd9bbfbae68cd523a0a34cae8fd01f3324f305affec8908508111614b0
3
+ size 342692
app/frontend/build/assets/index-D-qgc0vE.js ADDED
The diff for this file is too large to render. See raw diff
 
app/frontend/build/favicon.ico ADDED
app/frontend/build/fragmenta_background.png ADDED

Git LFS Details

  • SHA256: 048aea503935f9763e76db3f5d1fcd6d561d3db9aeac415605c46527a3d6631b
  • Pointer size: 131 Bytes
  • Size of remote file: 133 kB
app/frontend/build/fragmenta_icon_1024.png ADDED
app/frontend/build/index.html ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <!DOCTYPE html>
2
+ <html lang="en">
3
+ <head>
4
+ <meta charset="utf-8" />
5
+ <link rel="icon" href="/favicon.ico" />
6
+ <meta name="viewport" content="width=device-width, initial-scale=1" />
7
+ <meta name="theme-color" content="#000000" />
8
+ <meta
9
+ name="description"
10
+ content="Fragmenta Desktop - Stable Audio Fine-Tuning Application"
11
+ />
12
+ <link rel="manifest" href="/manifest.json" />
13
+
14
+ <link rel="preconnect" href="https://fonts.googleapis.com">
15
+ <link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>
16
+ <link href="https://fonts.googleapis.com/css2?family=JetBrains+Mono:wght@300;400;500;600&family=Space+Mono:wght@400;700&family=IBM+Plex+Mono:wght@300;400;500;600&display=swap" rel="stylesheet">
17
+
18
+ <style>
19
+ @font-face {
20
+ font-family: 'Bitcount Single';
21
+ src: url('/BitcountSingle-VariableFont_CRSV,ELSH,ELXP,slnt,wght.ttf') format('truetype');
22
+ font-weight: 100 900;
23
+ font-style: normal;
24
+ font-display: swap;
25
+ }
26
+ </style>
27
+
28
+ <title>Fragmenta Desktop</title>
29
+ <script type="module" crossorigin src="/assets/index-D-qgc0vE.js"></script>
30
+ </head>
31
+ <body>
32
+ <noscript>You need to enable JavaScript to run this app.</noscript>
33
+ <div id="root"></div>
34
+ </body>
35
+ </html>
app/frontend/build/manifest.json ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "short_name": "Fragmenta",
3
+ "name": "Fragmenta Desktop - Stable Audio Fine-Tuning",
4
+ "icons": [
5
+ {
6
+ "src": "favicon.ico",
7
+ "sizes": "64x64 32x32 24x24 16x16",
8
+ "type": "image/x-icon"
9
+ }
10
+ ],
11
+ "start_url": ".",
12
+ "display": "standalone",
13
+ "theme_color": "#000000",
14
+ "background_color": "#ffffff"
15
+ }
app/frontend/index.html ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <!DOCTYPE html>
2
+ <html lang="en">
3
+ <head>
4
+ <meta charset="utf-8" />
5
+ <link rel="icon" href="/favicon.ico" />
6
+ <meta name="viewport" content="width=device-width, initial-scale=1" />
7
+ <meta name="theme-color" content="#000000" />
8
+ <meta
9
+ name="description"
10
+ content="Fragmenta Desktop - Stable Audio Fine-Tuning Application"
11
+ />
12
+ <link rel="manifest" href="/manifest.json" />
13
+
14
+ <link rel="preconnect" href="https://fonts.googleapis.com">
15
+ <link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>
16
+ <link href="https://fonts.googleapis.com/css2?family=JetBrains+Mono:wght@300;400;500;600&family=Space+Mono:wght@400;700&family=IBM+Plex+Mono:wght@300;400;500;600&display=swap" rel="stylesheet">
17
+
18
+ <style>
19
+ @font-face {
20
+ font-family: 'Bitcount Single';
21
+ src: url('/BitcountSingle-VariableFont_CRSV,ELSH,ELXP,slnt,wght.ttf') format('truetype');
22
+ font-weight: 100 900;
23
+ font-style: normal;
24
+ font-display: swap;
25
+ }
26
+ </style>
27
+
28
+ <title>Fragmenta Desktop</title>
29
+ </head>
30
+ <body>
31
+ <noscript>You need to enable JavaScript to run this app.</noscript>
32
+ <div id="root"></div>
33
+ <script type="module" src="/src/index.js"></script>
34
+ </body>
35
+ </html>
app/frontend/package-lock.json CHANGED
The diff for this file is too large to render. See raw diff
 
app/frontend/package.json CHANGED
@@ -1,38 +1,23 @@
1
  {
2
  "name": "fragmenta-desktop",
3
- "version": "0.0.1",
4
  "description": "Fragmenta Desktop",
5
- "main": "index.js",
6
  "scripts": {
7
- "start": "react-scripts start",
8
- "build": "react-scripts build",
9
- "test": "react-scripts test",
10
- "eject": "react-scripts eject"
11
  },
12
  "dependencies": {
13
  "react": "^18.2.0",
14
  "react-dom": "^18.2.0",
15
- "react-scripts": "5.0.1",
16
- "axios": "^1.6.0",
17
  "@mui/material": "^5.14.0",
18
- "@mui/icons-material": "^5.14.0",
19
  "@emotion/react": "^11.11.0",
20
  "@emotion/styled": "^11.11.0",
21
- "react-dropzone": "^14.2.3",
22
- "react-player": "^2.13.0",
23
- "recharts": "^2.8.0"
24
  },
25
- "browserslist": {
26
- "production": [
27
- ">0.2%",
28
- "not dead",
29
- "not op_mini all"
30
- ],
31
- "development": [
32
- "last 1 chrome version",
33
- "last 1 firefox version",
34
- "last 1 safari version"
35
- ]
36
- },
37
- "proxy": "http://localhost:5001"
38
- }
 
1
  {
2
  "name": "fragmenta-desktop",
3
+ "version": "0.0.2",
4
  "description": "Fragmenta Desktop",
5
+ "type": "module",
6
  "scripts": {
7
+ "dev": "vite",
8
+ "build": "vite build",
9
+ "preview": "vite preview"
 
10
  },
11
  "dependencies": {
12
  "react": "^18.2.0",
13
  "react-dom": "^18.2.0",
 
 
14
  "@mui/material": "^5.14.0",
 
15
  "@emotion/react": "^11.11.0",
16
  "@emotion/styled": "^11.11.0",
17
+ "lucide-react": "^0.460.0"
 
 
18
  },
19
+ "devDependencies": {
20
+ "vite": "^5.4.0",
21
+ "@vitejs/plugin-react": "^4.3.0"
22
+ }
23
+ }
 
 
 
 
 
 
 
 
 
app/frontend/public/favicon.ico CHANGED

Git LFS Details

  • SHA256: 047116250e324c4c8b89804613534f3ba42073a7639e5d20bd06bcf923f60299
  • Pointer size: 131 Bytes
  • Size of remote file: 784 kB

Git LFS Details

  • SHA256: d858221015fbc58e25b68975d14751e707299f8a653d837205e0310e9294cee5
  • Pointer size: 129 Bytes
  • Size of remote file: 5.69 kB
app/frontend/public/fragmenta_icon_1024.png CHANGED

Git LFS Details

  • SHA256: 1e660c26e7b6a629a3addaafbb134586dd4d2d32f174d62edcc49d4be2ebc9af
  • Pointer size: 131 Bytes
  • Size of remote file: 845 kB

Git LFS Details

  • SHA256: 98e3ab5286c128bc8eaeac3b5ee55535218be5cc3d0fcbfde6b2fc2b68ae87ea
  • Pointer size: 130 Bytes
  • Size of remote file: 25.1 kB
app/frontend/src/App.js CHANGED
The diff for this file is too large to render. See raw diff
 
app/frontend/src/api.js ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ async function request(method, url, body, config = {}) {
2
+ const init = { method };
3
+ const headers = { ...(config.headers || {}) };
4
+
5
+ if (body !== undefined && body !== null) {
6
+ if (body instanceof FormData) {
7
+ init.body = body;
8
+ } else {
9
+ init.body = JSON.stringify(body);
10
+ if (!headers['Content-Type']) headers['Content-Type'] = 'application/json';
11
+ }
12
+ }
13
+ if (Object.keys(headers).length > 0) init.headers = headers;
14
+
15
+ const response = await fetch(url, init);
16
+
17
+ let data;
18
+ if (config.responseType === 'blob') {
19
+ data = await response.blob();
20
+ } else {
21
+ const text = await response.text();
22
+ try { data = text ? JSON.parse(text) : null; }
23
+ catch { data = text; }
24
+ }
25
+
26
+ if (!response.ok) {
27
+ const err = new Error(`HTTP ${response.status}`);
28
+ err.response = { status: response.status, data };
29
+ throw err;
30
+ }
31
+ return { data, status: response.status };
32
+ }
33
+
34
+ const api = {
35
+ get: (url, config) => request('GET', url, null, config),
36
+ post: (url, body, config) => request('POST', url, body, config),
37
+ put: (url, body, config) => request('PUT', url, body, config),
38
+ delete: (url, config) => request('DELETE', url, null, config),
39
+ };
40
+
41
+ export default api;
app/frontend/src/components/AudioUploadRow.js ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React, { useState, useEffect, useRef } from 'react';
2
+ import { Card, CardContent, Grid, Box, Typography, TextField, IconButton } from '@mui/material';
3
+ import { Upload as UploadIcon, Trash2 as DeleteIcon } from 'lucide-react';
4
+ import { audioUploadRowStyles } from '../theme';
5
+
6
+ export default function AudioUploadRow({ index, data, onChange, onRemove }) {
7
+ const [audioFile, setAudioFile] = useState(null);
8
+ const [audioUrl, setAudioUrl] = useState('');
9
+ const [isDragActive, setIsDragActive] = useState(false);
10
+ const inputRef = useRef(null);
11
+
12
+ useEffect(() => {
13
+ if (!data.file && !data.audioUrl) {
14
+ if (audioUrl) {
15
+ URL.revokeObjectURL(audioUrl);
16
+ }
17
+ setAudioFile(null);
18
+ setAudioUrl('');
19
+ }
20
+ }, [data.file, data.audioUrl, audioUrl]);
21
+
22
+ const acceptFile = (file) => {
23
+ if (!file || !file.type.startsWith('audio/')) return;
24
+ const url = URL.createObjectURL(file);
25
+ setAudioFile(file);
26
+ setAudioUrl(url);
27
+ onChange(index, { ...data, file, audioUrl: url });
28
+ };
29
+
30
+ return (
31
+ <Card sx={audioUploadRowStyles.card}>
32
+ <CardContent sx={audioUploadRowStyles.cardContent}>
33
+ <Grid container spacing={audioUploadRowStyles.gridSpacing} alignItems="center">
34
+ <Grid item xs={12} sm={4}>
35
+ <Box
36
+ onClick={() => inputRef.current?.click()}
37
+ onDragOver={(e) => { e.preventDefault(); setIsDragActive(true); }}
38
+ onDragLeave={() => setIsDragActive(false)}
39
+ onDrop={(e) => {
40
+ e.preventDefault();
41
+ setIsDragActive(false);
42
+ acceptFile(e.dataTransfer.files?.[0]);
43
+ }}
44
+ sx={audioUploadRowStyles.uploadDropZone(isDragActive)}
45
+ >
46
+ <input
47
+ ref={inputRef}
48
+ type="file"
49
+ accept="audio/*,.mp3,.wav,.flac,.m4a,.aac"
50
+ style={audioUploadRowStyles.hiddenInput}
51
+ onChange={(e) => acceptFile(e.target.files?.[0])}
52
+ />
53
+ {audioFile ? (
54
+ <Box>
55
+ <Typography variant="body2" color="textSecondary">
56
+ {audioFile.name}
57
+ </Typography>
58
+ {audioUrl && (
59
+ <audio
60
+ controls
61
+ src={audioUrl}
62
+ style={audioUploadRowStyles.audioPreview}
63
+ />
64
+ )}
65
+ </Box>
66
+ ) : (
67
+ <Box>
68
+ <UploadIcon size={20} color="#9198A1" />
69
+ <Typography variant="body2" color="textSecondary ">
70
+ {isDragActive ? 'Drop audio here' : ""}
71
+ </Typography>
72
+ </Box>
73
+ )}
74
+ </Box>
75
+ </Grid>
76
+
77
+ <Grid item xs={12} sm={7}>
78
+ <TextField
79
+ fullWidth
80
+ multiline
81
+ minRows={1}
82
+ maxRows={3}
83
+ label={`Prompt ${index + 1}`}
84
+ placeholder="Describe this audio file..."
85
+ value={data.prompt || ''}
86
+ onChange={(e) => onChange(index, { ...data, prompt: e.target.value })}
87
+ variant="outlined"
88
+ />
89
+ </Grid>
90
+
91
+ <Grid item xs={12} sm={1} sx={audioUploadRowStyles.deleteGridItem}>
92
+ <IconButton
93
+ color="error"
94
+ onClick={() => onRemove(index)}
95
+ sx={audioUploadRowStyles.deleteIconButton}
96
+ >
97
+ <DeleteIcon />
98
+ </IconButton>
99
+ </Grid>
100
+ </Grid>
101
+ </CardContent>
102
+ </Card>
103
+ );
104
+ }
app/frontend/src/components/BulkAnnotatePanel.js ADDED
@@ -0,0 +1,316 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React, { useState, useEffect, useRef, useCallback } from 'react';
2
+ import {
3
+ Box, Paper, Typography, TextField, Button, MenuItem, Select, FormControl,
4
+ InputLabel, LinearProgress, Alert, Table, TableHead, TableRow, TableCell,
5
+ TableBody, TableContainer, Checkbox, Tooltip, CircularProgress,
6
+ } from '@mui/material';
7
+ import {
8
+ Tags as TagsIcon,
9
+ CloudDownload as CloudDownloadIcon,
10
+ Save as SaveIcon,
11
+ FolderOpen as FolderOpenIcon,
12
+ } from 'lucide-react';
13
+ import api from '../api';
14
+
15
+ const POLL_INTERVAL_MS = 800;
16
+
17
+ export default function BulkAnnotatePanel({ onCommitted }) {
18
+ const [folderPath, setFolderPath] = useState('');
19
+ const [tier, setTier] = useState('basic');
20
+ const [status, setStatus] = useState(null);
21
+ const [results, setResults] = useState([]);
22
+ const [selected, setSelected] = useState({});
23
+ const [copyFiles, setCopyFiles] = useState(true);
24
+ const [message, setMessage] = useState('');
25
+ const [error, setError] = useState('');
26
+ const [committing, setCommitting] = useState(false);
27
+ const pollRef = useRef(null);
28
+
29
+ const stopPolling = useCallback(() => {
30
+ if (pollRef.current) {
31
+ clearInterval(pollRef.current);
32
+ pollRef.current = null;
33
+ }
34
+ }, []);
35
+
36
+ const fetchStatus = useCallback(async () => {
37
+ let data;
38
+ try {
39
+ const resp = await api.get('/api/bulk-annotate/status');
40
+ data = resp.data;
41
+ } catch (exc) {
42
+ // Transient errors (e.g. Flask auto-reload) must not kill polling —
43
+ // the download/annotation keeps running on the backend side.
44
+ return;
45
+ }
46
+ setStatus(data);
47
+
48
+ const annotationState = data.state;
49
+ const downloadState = data.clap_download?.state;
50
+
51
+ if (annotationState === 'done') {
52
+ try {
53
+ const resp = await api.get('/api/bulk-annotate/results');
54
+ setResults(resp.data.results || []);
55
+ const sel = {};
56
+ (resp.data.results || []).forEach((r, i) => { sel[i] = !r.error; });
57
+ setSelected(sel);
58
+ } catch {
59
+ return;
60
+ }
61
+ } else if (annotationState === 'error') {
62
+ setError(data.error || 'Annotation failed.');
63
+ }
64
+
65
+ const annotationInactive = annotationState !== 'running';
66
+ const downloadInactive = downloadState !== 'running';
67
+ if (annotationInactive && downloadInactive) {
68
+ stopPolling();
69
+ }
70
+ }, [stopPolling]);
71
+
72
+ const startPolling = useCallback(() => {
73
+ stopPolling();
74
+ pollRef.current = setInterval(fetchStatus, POLL_INTERVAL_MS);
75
+ }, [fetchStatus, stopPolling]);
76
+
77
+ useEffect(() => {
78
+ fetchStatus();
79
+ return () => stopPolling();
80
+ }, [fetchStatus, stopPolling]);
81
+
82
+ const startAnnotation = async () => {
83
+ setError('');
84
+ setMessage('');
85
+ setResults([]);
86
+ try {
87
+ await api.post('/api/bulk-annotate', { folder_path: folderPath, tier });
88
+ startPolling();
89
+ } catch (exc) {
90
+ setError(exc.response?.data?.error || exc.message);
91
+ }
92
+ };
93
+
94
+ const pickFolder = async () => {
95
+ setError('');
96
+ try {
97
+ const { data } = await api.post('/api/pick-folder', { start_dir: folderPath || undefined });
98
+ if (data?.path) setFolderPath(data.path);
99
+ } catch (exc) {
100
+ setError(exc.response?.data?.error || exc.message);
101
+ }
102
+ };
103
+
104
+ const downloadClap = async () => {
105
+ setError('');
106
+ try {
107
+ await api.post('/api/bulk-annotate/download-clap', {});
108
+ startPolling();
109
+ } catch (exc) {
110
+ setError(exc.response?.data?.error || exc.message);
111
+ }
112
+ };
113
+
114
+ const updatePrompt = (idx, value) => {
115
+ setResults(prev => prev.map((r, i) => (i === idx ? { ...r, prompt: value } : r)));
116
+ };
117
+
118
+ const toggleSelected = (idx) => {
119
+ setSelected(prev => ({ ...prev, [idx]: !prev[idx] }));
120
+ };
121
+
122
+ const toggleAll = () => {
123
+ const allSelected = results.every((_, i) => selected[i]);
124
+ const next = {};
125
+ results.forEach((_, i) => { next[i] = !allSelected; });
126
+ setSelected(next);
127
+ };
128
+
129
+ const commit = async () => {
130
+ setError('');
131
+ setMessage('');
132
+ setCommitting(true);
133
+ try {
134
+ const entries = results
135
+ .filter((_, i) => selected[i])
136
+ .map(r => ({ file_name: r.file_name, prompt: r.prompt, path: r.path }));
137
+ const { data } = await api.post('/api/bulk-annotate/commit', { entries, copy_files: copyFiles });
138
+ setMessage(data.message || 'Committed.');
139
+ if (onCommitted) onCommitted();
140
+ } catch (exc) {
141
+ setError(exc.response?.data?.error || exc.message);
142
+ } finally {
143
+ setCommitting(false);
144
+ }
145
+ };
146
+
147
+ const isRunning = status?.state === 'running';
148
+ const clapDownload = status?.clap_download;
149
+ const clapAvailable = !!status?.clap_available;
150
+ const clapDownloading = clapDownload?.state === 'running';
151
+ const richBlocked = tier === 'rich' && !clapAvailable;
152
+ const progressPct = status?.total ? Math.round((status.current / status.total) * 100) : 0;
153
+
154
+ return (
155
+ <Paper sx={{ p: 2, mt: 3 }} variant="outlined">
156
+ <Box sx={{ display: 'flex', alignItems: 'center', gap: 1, mb: 1.5 }}>
157
+ <TagsIcon size={20} />
158
+ <Typography variant="h6">Bulk Auto-Annotate</Typography>
159
+ </Box>
160
+ <Typography variant="body2" color="textSecondary" sx={{ mb: 2 }}>
161
+ Point at a folder of audio files and auto-generate prompts.
162
+ Basic uses librosa (tempo + key). Rich adds CLAP tagging (genre, mood, instruments).
163
+ </Typography>
164
+
165
+ <Box sx={{ display: 'flex', gap: 1.5, alignItems: 'center', flexWrap: 'wrap', mb: 2 }}>
166
+ <TextField
167
+ label="Folder path"
168
+ size="small"
169
+ value={folderPath}
170
+ onChange={(e) => setFolderPath(e.target.value)}
171
+ placeholder="Click Browse to choose a folder…"
172
+ sx={{ flexGrow: 1, minWidth: 260 }}
173
+ disabled={isRunning}
174
+ InputProps={{ readOnly: true }}
175
+ />
176
+ <Button
177
+ variant="outlined"
178
+ onClick={pickFolder}
179
+ startIcon={<FolderOpenIcon size={16} />}
180
+ disabled={isRunning}
181
+ >
182
+ Browse
183
+ </Button>
184
+ <FormControl size="small" sx={{ minWidth: 140 }} disabled={isRunning}>
185
+ <InputLabel id="tier-label">Tier</InputLabel>
186
+ <Select
187
+ labelId="tier-label"
188
+ value={tier}
189
+ label="Tier"
190
+ onChange={(e) => setTier(e.target.value)}
191
+ >
192
+ <MenuItem value="basic">Basic (no download)</MenuItem>
193
+ <MenuItem value="rich">Rich (CLAP, ~2.35 GB)</MenuItem>
194
+ </Select>
195
+ </FormControl>
196
+ <Tooltip title={richBlocked ? 'Download CLAP to enable Rich tier' : ''}>
197
+ <span>
198
+ <Button
199
+ variant="contained"
200
+ onClick={startAnnotation}
201
+ startIcon={isRunning ? <CircularProgress size={16} /> : <TagsIcon size={16} />}
202
+ disabled={isRunning || !folderPath || richBlocked}
203
+ >
204
+ {isRunning ? 'Annotating…' : 'Annotate'}
205
+ </Button>
206
+ </span>
207
+ </Tooltip>
208
+ {tier === 'rich' && !clapAvailable && (
209
+ <Button
210
+ variant="outlined"
211
+ onClick={downloadClap}
212
+ startIcon={clapDownloading ? <CircularProgress size={16} /> : <CloudDownloadIcon size={16} />}
213
+ disabled={clapDownloading}
214
+ >
215
+ {clapDownloading ? 'Downloading…' : 'Download CLAP'}
216
+ </Button>
217
+ )}
218
+ </Box>
219
+
220
+ {isRunning && (
221
+ <Box sx={{ mb: 2 }}>
222
+ <LinearProgress variant={status?.total ? 'determinate' : 'indeterminate'} value={progressPct} />
223
+ <Typography variant="caption" color="textSecondary">
224
+ {status?.current}/{status?.total} — {status?.current_file}
225
+ </Typography>
226
+ </Box>
227
+ )}
228
+ {clapDownload?.state === 'running' && (
229
+ <Alert severity="info" sx={{ mb: 2 }}>{clapDownload.message || 'Downloading CLAP checkpoint…'}</Alert>
230
+ )}
231
+ {clapDownload?.state === 'error' && (
232
+ <Alert severity="error" sx={{ mb: 2 }}>CLAP download failed: {clapDownload.error}</Alert>
233
+ )}
234
+ {error && <Alert severity="error" sx={{ mb: 2 }}>{error}</Alert>}
235
+ {message && <Alert severity="success" sx={{ mb: 2 }}>{message}</Alert>}
236
+
237
+ {results.length > 0 && (
238
+ <>
239
+ <TableContainer sx={{ maxHeight: 420, mb: 2 }}>
240
+ <Table size="small" stickyHeader>
241
+ <TableHead>
242
+ <TableRow>
243
+ <TableCell padding="checkbox">
244
+ <Checkbox
245
+ checked={results.every((_, i) => selected[i])}
246
+ indeterminate={
247
+ results.some((_, i) => selected[i]) && !results.every((_, i) => selected[i])
248
+ }
249
+ onChange={toggleAll}
250
+ />
251
+ </TableCell>
252
+ <TableCell>File</TableCell>
253
+ <TableCell>Prompt (editable)</TableCell>
254
+ </TableRow>
255
+ </TableHead>
256
+ <TableBody>
257
+ {results.map((row, idx) => (
258
+ <TableRow key={row.file_name + idx} hover>
259
+ <TableCell padding="checkbox">
260
+ <Checkbox
261
+ checked={!!selected[idx]}
262
+ onChange={() => toggleSelected(idx)}
263
+ disabled={!!row.error}
264
+ />
265
+ </TableCell>
266
+ <TableCell sx={{ maxWidth: 240, overflow: 'hidden', textOverflow: 'ellipsis', whiteSpace: 'nowrap' }}>
267
+ <Tooltip title={row.path || row.file_name}>
268
+ <span>{row.file_name}</span>
269
+ </Tooltip>
270
+ {row.error && (
271
+ <Typography variant="caption" color="error" display="block">
272
+ {row.error}
273
+ </Typography>
274
+ )}
275
+ </TableCell>
276
+ <TableCell>
277
+ <TextField
278
+ fullWidth
279
+ multiline
280
+ size="small"
281
+ value={row.prompt || ''}
282
+ onChange={(e) => updatePrompt(idx, e.target.value)}
283
+ disabled={!!row.error}
284
+ />
285
+ </TableCell>
286
+ </TableRow>
287
+ ))}
288
+ </TableBody>
289
+ </Table>
290
+ </TableContainer>
291
+
292
+ <Box sx={{ display: 'flex', gap: 2, alignItems: 'center' }}>
293
+ <FormControl size="small">
294
+ <Select
295
+ value={copyFiles ? 'copy' : 'link'}
296
+ onChange={(e) => setCopyFiles(e.target.value === 'copy')}
297
+ >
298
+ <MenuItem value="copy">Copy files into data/</MenuItem>
299
+ <MenuItem value="link">Leave files in place</MenuItem>
300
+ </Select>
301
+ </FormControl>
302
+ <Button
303
+ variant="contained"
304
+ color="primary"
305
+ onClick={commit}
306
+ startIcon={committing ? <CircularProgress size={16} /> : <SaveIcon size={16} />}
307
+ disabled={committing || !Object.values(selected).some(Boolean)}
308
+ >
309
+ {committing ? 'Saving…' : `Save ${Object.values(selected).filter(Boolean).length} to dataset`}
310
+ </Button>
311
+ </Box>
312
+ </>
313
+ )}
314
+ </Paper>
315
+ );
316
+ }
app/frontend/src/components/CheckpointManager.js ADDED
@@ -0,0 +1,211 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React, { useState } from 'react';
2
+ import {
3
+ Box,
4
+ Typography,
5
+ Card,
6
+ Chip,
7
+ Button,
8
+ Alert,
9
+ Dialog,
10
+ DialogTitle,
11
+ DialogContent,
12
+ DialogActions,
13
+ Snackbar,
14
+ } from '@mui/material';
15
+ import { CloudDownload as CloudDownloadIcon, Trash2 as DeleteIcon } from 'lucide-react';
16
+ import api from '../api';
17
+ import { checkpointManagerStyles } from '../theme';
18
+
19
+ export default function CheckpointManager({ model, onRefresh }) {
20
+ const [loadingStates, setLoadingStates] = useState({});
21
+ const [error, setError] = useState(null);
22
+ const [deleteTarget, setDeleteTarget] = useState(null);
23
+ const [toast, setToast] = useState({
24
+ open: false,
25
+ message: '',
26
+ severity: 'success'
27
+ });
28
+
29
+ const handleUnwrapCheckpoint = async (checkpoint) => {
30
+ const checkpointId = checkpoint.path;
31
+ setLoadingStates(prev => ({ ...prev, [checkpointId]: { unwrapping: true } }));
32
+ setError(null);
33
+ try {
34
+ await api.post('/api/unwrap-model', {
35
+ model_config: model.config_path,
36
+ ckpt_path: checkpoint.path,
37
+ name: `${checkpoint.name}_unwrapped`
38
+ });
39
+ setError(null);
40
+ setToast({
41
+ open: true,
42
+ message: `Checkpoint "${checkpoint.name}" unwrapped successfully.`,
43
+ severity: 'success'
44
+ });
45
+ onRefresh();
46
+ } catch (err) {
47
+ setError(`Failed to unwrap ${checkpoint.name}: ${err.response?.data?.error || err.message}`);
48
+ } finally {
49
+ setLoadingStates(prev => ({ ...prev, [checkpointId]: { unwrapping: false } }));
50
+ }
51
+ };
52
+
53
+ const handleDeleteCheckpoint = async () => {
54
+ if (!deleteTarget) {
55
+ return;
56
+ }
57
+
58
+ const checkpointId = deleteTarget.path;
59
+ setLoadingStates(prev => ({ ...prev, [checkpointId]: { deleting: true } }));
60
+ setError(null);
61
+
62
+ try {
63
+ await api.post('/api/delete-checkpoint', {
64
+ checkpoint_path: deleteTarget.path
65
+ });
66
+ setToast({
67
+ open: true,
68
+ message: `Checkpoint "${deleteTarget.name}" deleted successfully.`,
69
+ severity: 'success'
70
+ });
71
+ onRefresh();
72
+ } catch (err) {
73
+ setError(`Failed to delete ${deleteTarget.name}: ${err.response?.data?.error || err.message}`);
74
+ } finally {
75
+ setDeleteTarget(null);
76
+ setLoadingStates(prev => ({ ...prev, [checkpointId]: { deleting: false } }));
77
+ }
78
+ };
79
+
80
+ const closeToast = (_, reason) => {
81
+ if (reason === 'clickaway') {
82
+ return;
83
+ }
84
+ setToast(prev => ({ ...prev, open: false }));
85
+ };
86
+
87
+ const checkpoints = model.checkpoints || [];
88
+
89
+ return (
90
+ <>
91
+ <Box sx={checkpointManagerStyles.root}>
92
+ <Typography variant="subtitle2" color="textSecondary">
93
+ Checkpoints ({checkpoints.length})
94
+ </Typography>
95
+
96
+ {checkpoints.length === 0 ? (
97
+ <Typography variant="caption" color="textSecondary" sx={checkpointManagerStyles.emptyText}>
98
+ No checkpoints yet.
99
+ </Typography>
100
+ ) : (
101
+ <Box sx={checkpointManagerStyles.checkpointsList}>
102
+ {checkpoints.map((checkpoint, index) => {
103
+ const checkpointId = checkpoint.path;
104
+ const isUnwrapping = loadingStates[checkpointId]?.unwrapping;
105
+ const isDeleting = loadingStates[checkpointId]?.deleting;
106
+
107
+ const hasUnwrappedVersion = model.unwrapped_models?.some(unwrapped =>
108
+ unwrapped.name.includes(checkpoint.name) ||
109
+ checkpoint.name.includes(unwrapped.name.replace('_unwrapped', ''))
110
+ );
111
+
112
+ return (
113
+ <Card key={index} sx={checkpointManagerStyles.checkpointCard}>
114
+ <Box sx={checkpointManagerStyles.checkpointRow}>
115
+ <Box sx={checkpointManagerStyles.checkpointInfo}>
116
+ <Typography variant="body2" sx={checkpointManagerStyles.checkpointName}>
117
+ {checkpoint.name}
118
+ {hasUnwrappedVersion && (
119
+ <Chip
120
+ label="Unwrapped"
121
+ size="small"
122
+ color="success"
123
+ sx={checkpointManagerStyles.unwrappedChip}
124
+ />
125
+ )}
126
+ </Typography>
127
+ <Typography variant="caption" color="textSecondary">
128
+ {checkpoint.size_mb} MB
129
+ </Typography>
130
+ </Box>
131
+ <Box sx={checkpointManagerStyles.actions}>
132
+ {!hasUnwrappedVersion && (
133
+ <Button
134
+ variant="outlined"
135
+ color="primary"
136
+ size="small"
137
+ startIcon={<CloudDownloadIcon />}
138
+ onClick={() => handleUnwrapCheckpoint(checkpoint)}
139
+ disabled={isUnwrapping || isDeleting}
140
+ >
141
+ {isUnwrapping ? 'Unwrapping...' : 'Unwrap'}
142
+ </Button>
143
+ )}
144
+
145
+ {hasUnwrappedVersion && (
146
+ <Button
147
+ variant="outlined"
148
+ color="error"
149
+ size="small"
150
+ startIcon={<DeleteIcon />}
151
+ onClick={() => setDeleteTarget(checkpoint)}
152
+ disabled={isDeleting || isUnwrapping}
153
+ >
154
+ {isDeleting ? 'Deleting...' : 'Delete'}
155
+ </Button>
156
+ )}
157
+ </Box>
158
+ </Box>
159
+ </Card>
160
+ );
161
+ })}
162
+ </Box>
163
+ )}
164
+
165
+ {error && (
166
+ <Alert severity="error" sx={checkpointManagerStyles.errorAlert}>{error}</Alert>
167
+ )}
168
+ </Box>
169
+
170
+ <Dialog
171
+ open={Boolean(deleteTarget)}
172
+ onClose={() => setDeleteTarget(null)}
173
+ aria-labelledby="delete-checkpoint-dialog-title"
174
+ >
175
+ <DialogTitle id="delete-checkpoint-dialog-title">
176
+ Delete Wrapped Checkpoint
177
+ </DialogTitle>
178
+ <DialogContent>
179
+ <Typography sx={checkpointManagerStyles.deleteDialogText}>
180
+ {deleteTarget
181
+ ? `Delete "${deleteTarget.name}"? This action cannot be undone.`
182
+ : 'Delete this checkpoint?'}
183
+ </Typography>
184
+ </DialogContent>
185
+ <DialogActions>
186
+ <Button onClick={() => setDeleteTarget(null)}>
187
+ Cancel
188
+ </Button>
189
+ <Button
190
+ variant="contained"
191
+ color="error"
192
+ onClick={handleDeleteCheckpoint}
193
+ >
194
+ Delete
195
+ </Button>
196
+ </DialogActions>
197
+ </Dialog>
198
+
199
+ <Snackbar
200
+ open={toast.open}
201
+ autoHideDuration={3200}
202
+ onClose={closeToast}
203
+ anchorOrigin={{ vertical: 'bottom', horizontal: 'right' }}
204
+ >
205
+ <Alert onClose={closeToast} severity={toast.severity} variant="filled" sx={checkpointManagerStyles.snackbarAlert}>
206
+ {toast.message}
207
+ </Alert>
208
+ </Snackbar>
209
+ </>
210
+ );
211
+ }
app/frontend/src/components/GeneratedFragmentsWindow.js ADDED
@@ -0,0 +1,114 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React, { useState, useRef, useCallback } from 'react';
2
+ import { Paper, Box, Typography, Button, List, ListItem, IconButton } from '@mui/material';
3
+ import { Square as StopIcon, Play as PlayIcon, Download as DownloadIcon } from 'lucide-react';
4
+ import api from '../api';
5
+ import { generatedFragmentsWindowStyles } from '../theme';
6
+
7
+ export default function GeneratedFragmentsWindow({ fragments, onDownload }) {
8
+ const [playingFragment, setPlayingFragment] = useState(null);
9
+ const audioRefs = useRef({});
10
+
11
+ const handlePlayPause = (fragment) => {
12
+ const audio = audioRefs.current[fragment.id];
13
+ if (!audio) return;
14
+
15
+ if (playingFragment === fragment.id) {
16
+ audio.pause();
17
+ setPlayingFragment(null);
18
+ } else {
19
+ if (playingFragment && audioRefs.current[playingFragment]) {
20
+ audioRefs.current[playingFragment].pause();
21
+ }
22
+ audio.play();
23
+ setPlayingFragment(fragment.id);
24
+ }
25
+ };
26
+
27
+ const setAudioRef = useCallback((fragmentId, audioElement) => {
28
+ if (audioElement) {
29
+ audioRefs.current[fragmentId] = audioElement;
30
+ }
31
+ }, []);
32
+
33
+ return (
34
+ <Paper
35
+ variant="outlined"
36
+ sx={generatedFragmentsWindowStyles.rootPaper}
37
+ >
38
+ <Box sx={generatedFragmentsWindowStyles.headerRow}>
39
+ <Box sx={generatedFragmentsWindowStyles.titleRow}>
40
+ <Box component="span" sx={generatedFragmentsWindowStyles.titleIcon}>
41
+ <DownloadIcon size={20} />
42
+ </Box>
43
+ <Typography variant="h6" sx={generatedFragmentsWindowStyles.titleText}>
44
+ Generated Fragments
45
+ </Typography>
46
+ </Box>
47
+ <Typography variant="caption" color="textSecondary" sx={generatedFragmentsWindowStyles.countText}>
48
+ {fragments.length}
49
+ </Typography>
50
+ </Box>
51
+
52
+ {fragments.length === 0 ? (
53
+ <Box
54
+ sx={generatedFragmentsWindowStyles.emptyState}
55
+ >
56
+ <Typography variant="body2">
57
+ No fragments generated yet
58
+ </Typography>
59
+ </Box>
60
+ ) : (
61
+ <List
62
+ sx={generatedFragmentsWindowStyles.listRoot}
63
+ >
64
+ {fragments.slice().reverse().map((fragment) => (
65
+ <ListItem
66
+ key={fragment.id}
67
+ sx={generatedFragmentsWindowStyles.listItem}
68
+ >
69
+ <Box sx={generatedFragmentsWindowStyles.fragmentRow}>
70
+ <Box sx={generatedFragmentsWindowStyles.fragmentMeta}>
71
+ <Typography
72
+ variant="subtitle2"
73
+ sx={generatedFragmentsWindowStyles.fragmentPrompt}
74
+ >
75
+ {fragment.prompt}
76
+ </Typography>
77
+ <Typography variant="caption" color="textSecondary">
78
+ {fragment.duration}s • {fragment.timestamp}
79
+ </Typography>
80
+ </Box>
81
+ <Box sx={generatedFragmentsWindowStyles.fragmentActions}>
82
+ <IconButton
83
+ size="small"
84
+ onClick={() => handlePlayPause(fragment)}
85
+ color={playingFragment === fragment.id ? "primary" : "default"}
86
+ sx={generatedFragmentsWindowStyles.playPauseButton(playingFragment === fragment.id)}
87
+ >
88
+ {playingFragment === fragment.id ? <StopIcon /> : <PlayIcon />}
89
+ </IconButton>
90
+ <Button
91
+ size="small"
92
+ variant="outlined"
93
+ startIcon={<DownloadIcon />}
94
+ onClick={() => onDownload(fragment)}
95
+ >
96
+ Download
97
+ </Button>
98
+ </Box>
99
+ </Box>
100
+
101
+ <audio
102
+ ref={el => setAudioRef(fragment.id, el)}
103
+ src={fragment.audioUrl}
104
+ onEnded={() => setPlayingFragment(null)}
105
+ onPause={() => setPlayingFragment(null)}
106
+ style={generatedFragmentsWindowStyles.hiddenAudio}
107
+ />
108
+ </ListItem>
109
+ ))}
110
+ </List>
111
+ )}
112
+ </Paper>
113
+ );
114
+ }
app/frontend/src/components/HfAuthDialog.js ADDED
@@ -0,0 +1,245 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React, { useState, useEffect } from 'react';
2
+ import {
3
+ Dialog,
4
+ DialogTitle,
5
+ DialogContent,
6
+ DialogActions,
7
+ Button,
8
+ Typography,
9
+ TextField,
10
+ Box,
11
+ CircularProgress,
12
+ Stepper,
13
+ Step,
14
+ StepLabel,
15
+ Link,
16
+ Alert,
17
+ LinearProgress
18
+ } from '@mui/material';
19
+ import api from '../api';
20
+ import { hfAuthDialogStyles } from '../theme';
21
+
22
+ const HfAuthDialog = ({ open, onClose, onModelsDownloaded }) => {
23
+ const [activeStep, setActiveStep] = useState(0);
24
+ const [missingModels, setMissingModels] = useState([]);
25
+ const [checkingStatus, setCheckingStatus] = useState(true);
26
+ const [token, setToken] = useState('');
27
+ const [error, setError] = useState(null);
28
+ const [isProcessing, setIsProcessing] = useState(false);
29
+ const [downloadingModel, setDownloadingModel] = useState(null);
30
+
31
+ const steps = ['Check required models', 'Authenticate', 'Download models'];
32
+
33
+ useEffect(() => {
34
+ if (open) {
35
+ checkModelStatus();
36
+ } else {
37
+ // Reset state on close
38
+ setActiveStep(0);
39
+ setError(null);
40
+ setToken('');
41
+ setMissingModels([]);
42
+ }
43
+ }, [open]);
44
+
45
+ const checkModelStatus = async () => {
46
+ setCheckingStatus(true);
47
+ setError(null);
48
+ try {
49
+ const response = await api.get('/api/base-models/status');
50
+ const models = response.data.base_models;
51
+ const missing = Object.entries(models)
52
+ .filter(([_, info]) => !info.downloaded)
53
+ .map(([id, info]) => ({ id, ...info }));
54
+
55
+ setMissingModels(missing);
56
+
57
+ if (missing.length === 0) {
58
+ // All models exist
59
+ setActiveStep(3);
60
+ } else {
61
+ setActiveStep(1);
62
+ }
63
+ } catch (err) {
64
+ setError(err.response?.data?.error || err.message || 'Failed to check model status.');
65
+ } finally {
66
+ setCheckingStatus(false);
67
+ }
68
+ };
69
+
70
+ const handleLogin = async () => {
71
+ if (!token.trim()) {
72
+ setError('Please enter a Hugging Face token.');
73
+ return;
74
+ }
75
+
76
+ setIsProcessing(true);
77
+ setError(null);
78
+ try {
79
+ await api.post('/api/hf-login', { token: token.trim() });
80
+
81
+ // If login successful, move to download
82
+ setActiveStep(2);
83
+ startDownloads();
84
+ } catch (err) {
85
+ setError(err.response?.data?.error || err.message || 'Authentication failed. Please check your token.');
86
+ setIsProcessing(false);
87
+ }
88
+ };
89
+
90
+ const startDownloads = async () => {
91
+ try {
92
+ for (const model of missingModels) {
93
+ setDownloadingModel(model.name);
94
+
95
+ // Record terms acceptance
96
+ await api.post(`/api/models/${model.id}/accept-terms`);
97
+
98
+ await api.post(`/api/models/${model.id}/download`);
99
+ }
100
+
101
+ // All done
102
+ setActiveStep(3);
103
+ if (onModelsDownloaded) {
104
+ onModelsDownloaded();
105
+ }
106
+ } catch (err) {
107
+ setError(err.response?.data?.error || err.message || 'Failed to download models.');
108
+ } finally {
109
+ setIsProcessing(false);
110
+ setDownloadingModel(null);
111
+ }
112
+ };
113
+
114
+ const handleClose = () => {
115
+ if (isProcessing && activeStep === 2) {
116
+ // Cannot close while downloading
117
+ return;
118
+ }
119
+ onClose(activeStep === 3); // return true if finished successfully
120
+ };
121
+
122
+ const getStepContent = (stepIndex) => {
123
+ if (checkingStatus) {
124
+ return (
125
+ <Box sx={hfAuthDialogStyles.checkingBox}>
126
+ <CircularProgress sx={hfAuthDialogStyles.checkingProgress} />
127
+ <Typography>Checking model availability...</Typography>
128
+ </Box>
129
+ );
130
+ }
131
+
132
+ switch (stepIndex) {
133
+ case 1:
134
+ return (
135
+ <Box sx={hfAuthDialogStyles.authStepBox}>
136
+ <Typography variant="body1" paragraph>
137
+ Some required base models are missing. You need a Hugging Face access token to download them.
138
+ </Typography>
139
+
140
+ <Typography variant="body2" color="textSecondary" paragraph>
141
+ Before proceeding, please ensure you have visited huggingface.co and accepted the terms of use for the required models (e.g., stabilityai/stable-audio-open-1.0).
142
+ </Typography>
143
+
144
+ <TextField
145
+ fullWidth
146
+ label="Hugging Face Access Token"
147
+ type="password"
148
+ value={token}
149
+ onChange={(e) => setToken(e.target.value)}
150
+ placeholder="hf_xxxxxxxxxxxxxxxxxxx..."
151
+ margin="normal"
152
+ variant="outlined"
153
+ />
154
+ <Typography variant="caption" color="textSecondary">
155
+ You can get a token from{' '}
156
+ <Link href="https://huggingface.co/settings/tokens" target="_blank" rel="noopener">
157
+ your Hugging Face settings
158
+ </Link>. "Read" access to public gated repos is needed.
159
+ </Typography>
160
+ </Box>
161
+ );
162
+ case 2:
163
+ return (
164
+ <Box sx={hfAuthDialogStyles.downloadStepBox}>
165
+ <Typography variant="h6" paragraph>
166
+ Downloading {downloadingModel}...
167
+ </Typography>
168
+ <LinearProgress sx={hfAuthDialogStyles.downloadProgress} />
169
+ <Typography variant="body2" color="textSecondary">
170
+ This may take several minutes depending on your connection speed. Do not close the application.
171
+ </Typography>
172
+ </Box>
173
+ );
174
+ case 3:
175
+ return (
176
+ <Box sx={hfAuthDialogStyles.successStepBox}>
177
+ <Typography variant="h6" color="success.main" paragraph>
178
+ All models are ready!
179
+ </Typography>
180
+ <Typography variant="body1">
181
+ You can now close this dialog and begin using Fragmenta.
182
+ </Typography>
183
+ </Box>
184
+ );
185
+ default:
186
+ return "Unknown step";
187
+ }
188
+ };
189
+
190
+ return (
191
+ <Dialog
192
+ open={open}
193
+ onClose={handleClose}
194
+ maxWidth="sm"
195
+ fullWidth
196
+ disableEscapeKeyDown={isProcessing && activeStep === 2}
197
+ >
198
+ <DialogTitle>
199
+ Hugging Face Authentication
200
+ </DialogTitle>
201
+ <DialogContent dividers>
202
+ <Stepper activeStep={activeStep} alternativeLabel sx={hfAuthDialogStyles.stepper}>
203
+ {steps.map((label) => (
204
+ <Step key={label}>
205
+ <StepLabel>{label}</StepLabel>
206
+ </Step>
207
+ ))}
208
+ </Stepper>
209
+
210
+ {error && (
211
+ <Alert severity="error" sx={hfAuthDialogStyles.errorAlert}>{error}</Alert>
212
+ )}
213
+
214
+ {getStepContent(activeStep)}
215
+ </DialogContent>
216
+
217
+ <DialogActions>
218
+ {activeStep !== 3 && activeStep !== 2 && (
219
+ <Button onClick={handleClose} disabled={isProcessing}>
220
+ Cancel
221
+ </Button>
222
+ )}
223
+
224
+ {activeStep === 1 && (
225
+ <Button
226
+ variant="contained"
227
+ color="primary"
228
+ onClick={handleLogin}
229
+ disabled={isProcessing || !token}
230
+ >
231
+ {isProcessing ? <CircularProgress size={hfAuthDialogStyles.loginSpinnerSize} /> : 'Login & Download'}
232
+ </Button>
233
+ )}
234
+
235
+ {activeStep === 3 && (
236
+ <Button variant="contained" color="primary" onClick={handleClose}>
237
+ Close
238
+ </Button>
239
+ )}
240
+ </DialogActions>
241
+ </Dialog>
242
+ );
243
+ };
244
+
245
+ export default HfAuthDialog;
app/frontend/src/components/LossChart.js ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React, { useState } from 'react';
2
+ import { lossChartStyles } from '../theme';
3
+
4
+ function fmtTime(sec) {
5
+ const m = Math.floor(sec / 60);
6
+ const s = Math.floor(sec % 60).toString().padStart(2, '0');
7
+ return `${m}:${s}`;
8
+ }
9
+
10
+ export default function LossChart({ data, width = 600, height = 200 }) {
11
+ const [hover, setHover] = useState(null);
12
+ const padding = lossChartStyles.padding;
13
+ const colors = lossChartStyles.colors;
14
+ const axisFontSize = lossChartStyles.axisFontSize;
15
+ const tooltip = lossChartStyles.tooltip;
16
+
17
+ if (!data || data.length === 0) return null;
18
+
19
+ const innerW = width - padding.left - padding.right;
20
+ const innerH = height - padding.top - padding.bottom;
21
+
22
+ const xs = data.map(d => d.time);
23
+ const ys = data.map(d => d.loss);
24
+ const xMin = Math.min(...xs);
25
+ const xMax = Math.max(...xs);
26
+ const yMin = Math.min(...ys);
27
+ const yMax = Math.max(...ys);
28
+ const xRange = xMax - xMin || 1;
29
+ const yRange = yMax - yMin || 1;
30
+
31
+ const xScale = v => padding.left + ((v - xMin) / xRange) * innerW;
32
+ const yScale = v => padding.top + innerH - ((v - yMin) / yRange) * innerH;
33
+
34
+ const points = data.map(d => `${xScale(d.time)},${yScale(d.loss)}`).join(' ');
35
+
36
+ const yTicks = 4;
37
+ const xTicks = Math.min(5, data.length);
38
+
39
+ const handleMove = (e) => {
40
+ const rect = e.currentTarget.getBoundingClientRect();
41
+ const px = ((e.clientX - rect.left) / rect.width) * width;
42
+ let nearest = data[0];
43
+ let bestDist = Infinity;
44
+ for (const d of data) {
45
+ const dist = Math.abs(xScale(d.time) - px);
46
+ if (dist < bestDist) { bestDist = dist; nearest = d; }
47
+ }
48
+ setHover(nearest);
49
+ };
50
+
51
+ return (
52
+ <svg
53
+ viewBox={`0 0 ${width} ${height}`}
54
+ preserveAspectRatio="none"
55
+ style={lossChartStyles.svg}
56
+ onMouseMove={handleMove}
57
+ onMouseLeave={() => setHover(null)}
58
+ >
59
+ {Array.from({ length: yTicks + 1 }, (_, i) => {
60
+ const v = yMin + (yRange * i) / yTicks;
61
+ const y = yScale(v);
62
+ return (
63
+ <g key={`y${i}`}>
64
+ <line x1={padding.left} x2={width - padding.right} y1={y} y2={y}
65
+ stroke={colors.grid} strokeDasharray="3 3" />
66
+ <text x={padding.left - 6} y={y + 4} textAnchor="end"
67
+ fontSize={axisFontSize} fill={colors.axis}>{v.toFixed(3)}</text>
68
+ </g>
69
+ );
70
+ })}
71
+
72
+ {Array.from({ length: xTicks }, (_, i) => {
73
+ const v = xMin + (xRange * i) / Math.max(xTicks - 1, 1);
74
+ const x = xScale(v);
75
+ return (
76
+ <text key={`x${i}`} x={x} y={height - padding.bottom + 16}
77
+ textAnchor="middle" fontSize={axisFontSize} fill={colors.axis}>
78
+ {fmtTime(v)}
79
+ </text>
80
+ );
81
+ })}
82
+
83
+ <polyline fill="none" stroke={colors.line} strokeWidth="2" points={points} />
84
+
85
+ {data.map((d, i) => (
86
+ <circle key={i} cx={xScale(d.time)} cy={yScale(d.loss)} r="2" fill={colors.point} />
87
+ ))}
88
+
89
+ {hover && (
90
+ <g>
91
+ <line x1={xScale(hover.time)} x2={xScale(hover.time)}
92
+ y1={padding.top} y2={height - padding.bottom}
93
+ stroke={colors.axis} strokeDasharray="2 2" />
94
+ <circle cx={xScale(hover.time)} cy={yScale(hover.loss)} r="4" fill={colors.line} />
95
+ <g transform={`translate(${Math.min(xScale(hover.time) + 8, width - (tooltip.width + 10))}, ${padding.top + 6})`}>
96
+ <rect width={tooltip.width} height={tooltip.height} rx={tooltip.rx} fill={colors.tooltipBg} stroke={colors.tooltipBorder} />
97
+ <text x={tooltip.textX} y={tooltip.timeY} fontSize={axisFontSize} fill={colors.tooltipText}>Time: {fmtTime(hover.time)}</text>
98
+ <text x={tooltip.textX} y={tooltip.lossY} fontSize={axisFontSize} fill={colors.tooltipText}>Loss: {hover.loss.toFixed(4)}</text>
99
+ </g>
100
+ </g>
101
+ )}
102
+ </svg>
103
+ );
104
+ }
app/frontend/src/components/ModelUnwrapButton.js ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React, { useState } from 'react';
2
+ import { Button, Box, Link, Typography } from '@mui/material';
3
+ import { CloudDownload as CloudDownloadIcon } from 'lucide-react';
4
+ import api from '../api';
5
+ import { modelUnwrapButtonStyles } from '../theme';
6
+
7
+ export default function ModelUnwrapButton({ model, onUnwrap, onRefresh }) {
8
+ const [loading, setLoading] = useState(false);
9
+ const [result, setResult] = useState(null);
10
+ const [error, setError] = useState(null);
11
+
12
+ const handleUnwrap = async () => {
13
+ setLoading(true);
14
+ setResult(null);
15
+ setError(null);
16
+
17
+ try {
18
+ const response = await api.post('/api/unwrap-model', {
19
+ model_config: model.configPath,
20
+ ckpt_path: model.ckptPath,
21
+ name: model.name + '_unwrapped'
22
+ });
23
+ setResult(response.data);
24
+ if (onUnwrap) onUnwrap(response.data);
25
+ if (onRefresh) onRefresh();
26
+ } catch (err) {
27
+ console.error('Unwrap error:', err);
28
+ setError(err.response?.data?.error || err.message);
29
+ } finally {
30
+ setLoading(false);
31
+ }
32
+ };
33
+
34
+ return (
35
+ <Box sx={modelUnwrapButtonStyles.root}>
36
+ <Button
37
+ variant="outlined"
38
+ color="primary"
39
+ size="small"
40
+ startIcon={<CloudDownloadIcon />}
41
+ onClick={handleUnwrap}
42
+ disabled={loading}
43
+ >
44
+ {loading ? 'Unwrapping...' : 'Unwrap for Inference'}
45
+ </Button>
46
+ {result && result.unwrapped_path && (
47
+ <Box sx={modelUnwrapButtonStyles.result}>
48
+ <Link href={`file://${result.unwrapped_path}`} target="_blank" rel="noopener noreferrer">
49
+ Download Unwrapped Model
50
+ </Link>
51
+ </Box>
52
+ )}
53
+ {error && (
54
+ <Typography sx={modelUnwrapButtonStyles.error}>{error}</Typography>
55
+ )}
56
+ </Box>
57
+ );
58
+ }
app/frontend/src/components/TabPanel.js ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ 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>
18
+ )}
19
+ </div>
20
+ );
21
+ }
app/frontend/src/components/TrainingMonitor.js ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React from 'react';
2
+ import { Paper, Box, Typography, LinearProgress, Grid, Alert } from '@mui/material';
3
+ import { Activity as ActivityIcon } from 'lucide-react';
4
+ import LossChart from './LossChart';
5
+ import { trainingMonitorStyles } from '../theme';
6
+
7
+ export default function TrainingMonitor({
8
+ trainingProgress,
9
+ trainingStatus,
10
+ trainingError,
11
+ trainingConfig,
12
+ indicatorState,
13
+ }) {
14
+ const getProgressColor = () => {
15
+ if (trainingError) return 'error';
16
+ if (trainingProgress === 100) return 'success';
17
+ return 'primary';
18
+ };
19
+
20
+ const status = indicatorState?.status || 'idle';
21
+ const label = indicatorState?.label || 'Idle';
22
+ const animate = indicatorState?.animate || false;
23
+
24
+ return (
25
+ <Paper sx={trainingMonitorStyles.rootPaper}>
26
+ <Box sx={trainingMonitorStyles.headerRow}>
27
+ <Box sx={trainingMonitorStyles.headerTitleWrap}>
28
+ <Box component="span" sx={trainingMonitorStyles.headerIcon}>
29
+ <ActivityIcon size={20} />
30
+ </Box>
31
+ <Typography variant="h6" sx={trainingMonitorStyles.headerTitle}>
32
+ Training Monitor
33
+ </Typography>
34
+ </Box>
35
+ <Box sx={trainingMonitorStyles.statusInline}>
36
+ <Box sx={trainingMonitorStyles.statusDot(status, animate)} />
37
+ <Typography variant="caption" sx={trainingMonitorStyles.statusText(status)}>
38
+ {label}
39
+ </Typography>
40
+ </Box>
41
+ </Box>
42
+
43
+ <Box sx={trainingMonitorStyles.progressSection}>
44
+ <Box sx={trainingMonitorStyles.progressHeader}>
45
+ <Typography variant="body2">Progress</Typography>
46
+ <Typography variant="body2">{trainingProgress}%</Typography>
47
+ </Box>
48
+ <LinearProgress
49
+ variant="determinate"
50
+ value={trainingProgress}
51
+ color={getProgressColor()}
52
+ sx={trainingMonitorStyles.progressBar}
53
+ />
54
+ </Box>
55
+
56
+ {trainingStatus?.device_info && (
57
+ <Box sx={trainingMonitorStyles.deviceSection}>
58
+ <Typography variant="body2" color="textSecondary" gutterBottom>
59
+ <strong>Device Used for Training</strong>
60
+ </Typography>
61
+ <Typography variant="body2">
62
+ Device: {trainingStatus.device_info.device} ({trainingStatus.device_info.memory_gb?.toFixed(2)}GB VRAM)
63
+ </Typography>
64
+ <Typography variant="body2" color="textSecondary" sx={trainingMonitorStyles.deviceInfo}>
65
+ Info: {trainingStatus.device_info.type === 'cuda' ? 'CUDA GPU available and selected for training' :
66
+ trainingStatus.device_info.type === 'cpu' ? 'Using CPU (no CUDA GPU available or compatible)' :
67
+ 'Using MPS (Apple Silicon GPU)'}
68
+ </Typography>
69
+ </Box>
70
+ )}
71
+
72
+ <Grid container spacing={2} sx={trainingMonitorStyles.metricsGrid}>
73
+ <Grid item xs={12} sm={6}>
74
+ <Typography variant="body2" color="textSecondary">Current Epoch</Typography>
75
+ <Typography variant="body1">
76
+ {trainingStatus?.current_epoch !== undefined ?
77
+ `${trainingStatus.current_epoch + 1} / ${trainingConfig.epochs}` :
78
+ '0 / ' + trainingConfig.epochs}
79
+ </Typography>
80
+ </Grid>
81
+ <Grid item xs={12} sm={6}>
82
+ <Typography variant="body2" color="textSecondary">Global Step / Total Steps</Typography>
83
+ <Typography variant="body1" color="primary">
84
+ {trainingStatus?.global_step !== undefined && trainingStatus?.total_steps !== undefined ?
85
+ `${trainingStatus.global_step} / ${trainingStatus.total_steps}` :
86
+ 'N/A'}
87
+ </Typography>
88
+ </Grid>
89
+ <Grid item xs={12} sm={6}>
90
+ <Typography variant="body2" color="textSecondary">Checkpoints Saved</Typography>
91
+ <Typography variant="body1">
92
+ {trainingStatus?.checkpoints_saved || 0}
93
+ </Typography>
94
+ </Grid>
95
+ <Grid item xs={12} sm={6}>
96
+ <Typography variant="body2" color="textSecondary">Current Loss</Typography>
97
+ <Typography variant="body1">
98
+ {trainingStatus?.loss ? parseFloat(trainingStatus.loss).toFixed(4) : 'N/A'}
99
+ </Typography>
100
+ </Grid>
101
+ </Grid>
102
+
103
+ {trainingStatus?.loss_history && trainingStatus.loss_history.length > 0 && (
104
+ <Box sx={trainingMonitorStyles.lossSection}>
105
+ <Typography variant="body2" color="textSecondary" gutterBottom>
106
+ <strong>Loss History</strong>
107
+ </Typography>
108
+ <Box sx={trainingMonitorStyles.lossChartBox}>
109
+ <LossChart data={trainingStatus.loss_history} />
110
+ </Box>
111
+ </Box>
112
+ )}
113
+
114
+ {trainingError && (
115
+ <Alert severity="error" sx={trainingMonitorStyles.errorAlert}>
116
+ <Typography variant="body2">
117
+ <strong>Training Error:</strong> {trainingError}
118
+ </Typography>
119
+ </Alert>
120
+ )}
121
+ </Paper>
122
+ );
123
+ }
app/frontend/src/components/WelcomePage.js ADDED
@@ -0,0 +1,102 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React, { useState, useEffect } from 'react';
2
+ import { Backdrop, Box, Fade, Typography, Button, Checkbox, FormControlLabel } from '@mui/material';
3
+ import { welcomePageStyles } from '../theme';
4
+
5
+ export default function WelcomePage({ open, onClose }) {
6
+ const [titleVisible, setTitleVisible] = useState(false);
7
+ const [textVisible, setTextVisible] = useState(false);
8
+ const [dontShowAgain, setDontShowAgain] = useState(false);
9
+
10
+ useEffect(() => {
11
+ if (open) {
12
+ const titleTimer = setTimeout(() => setTitleVisible(true), 500);
13
+ const textTimer = setTimeout(() => setTextVisible(true), 1500);
14
+ return () => {
15
+ clearTimeout(titleTimer);
16
+ clearTimeout(textTimer);
17
+ };
18
+ } else {
19
+ setTitleVisible(false);
20
+ setTextVisible(false);
21
+ }
22
+ }, [open]);
23
+
24
+ if (!open) return null;
25
+
26
+ return (
27
+ <Backdrop
28
+ open={open}
29
+ onClick={() => onClose(false)}
30
+ sx={welcomePageStyles.backdrop}
31
+ >
32
+ <Box
33
+ sx={welcomePageStyles.panel}
34
+ onClick={(e) => e.stopPropagation()}
35
+ >
36
+ <Fade in={titleVisible} timeout={800}>
37
+ <Box sx={welcomePageStyles.logo} />
38
+ </Fade>
39
+
40
+ <Fade in={titleVisible} timeout={1000}>
41
+ <Typography
42
+ variant="h2"
43
+ component="h1"
44
+ sx={welcomePageStyles.title}
45
+ >
46
+ Welcome to Fragmenta
47
+ </Typography>
48
+ </Fade>
49
+
50
+ <Fade in={textVisible} timeout={1000}>
51
+ <Box>
52
+ <Typography
53
+ variant="overline"
54
+ sx={welcomePageStyles.overline}
55
+ >
56
+ An End-to-End Pipeline to Fine-Tune and Use Text-to-Audio Models.
57
+ </Typography>
58
+
59
+
60
+ <Typography
61
+ variant="body2"
62
+ sx={welcomePageStyles.footer}
63
+ >
64
+ @2025-2026 Misagh Azimi
65
+ </Typography>
66
+ <Typography
67
+ variant="body2"
68
+ sx={welcomePageStyles.version}
69
+ >
70
+ Version 0.0.2
71
+ </Typography>
72
+ <Button
73
+ variant="contained"
74
+ onClick={() => onClose(dontShowAgain)}
75
+ sx={welcomePageStyles.ctaButton}
76
+ >
77
+ Get Started
78
+ </Button>
79
+ <Box sx={{ mt: 1.5 }}>
80
+ <FormControlLabel
81
+ control={
82
+ <Checkbox
83
+ checked={dontShowAgain}
84
+ onChange={(e) => setDontShowAgain(e.target.checked)}
85
+ size="small"
86
+ sx={{ color: 'text.secondary' }}
87
+ />
88
+ }
89
+ label={
90
+ <Typography variant="caption" sx={{ color: 'text.secondary' }}>
91
+ Don't show this again
92
+ </Typography>
93
+ }
94
+ />
95
+ </Box>
96
+
97
+ </Box>
98
+ </Fade>
99
+ </Box>
100
+ </Backdrop>
101
+ );
102
+ }
app/frontend/src/index.js CHANGED
@@ -4,16 +4,8 @@ import App from './App';
4
 
5
  document.body.style.margin = '0';
6
  document.body.style.padding = '0';
7
- document.body.style.backgroundColor = '#0D1117';
8
-
9
- document.body.style.overflow = 'auto';
10
- document.documentElement.style.backgroundColor = '#0D1117';
11
-
12
- document.documentElement.style.overflow = 'auto';
13
 
14
  const rootElement = document.getElementById('root');
15
-
16
- rootElement.style.overflow = 'auto';
17
  rootElement.style.minHeight = '100vh';
18
 
19
  const root = ReactDOM.createRoot(rootElement);
 
4
 
5
  document.body.style.margin = '0';
6
  document.body.style.padding = '0';
 
 
 
 
 
 
7
 
8
  const rootElement = document.getElementById('root');
 
 
9
  rootElement.style.minHeight = '100vh';
10
 
11
  const root = ReactDOM.createRoot(rootElement);
app/frontend/src/theme.js ADDED
@@ -0,0 +1,2113 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { createTheme, responsiveFontSizes } from '@mui/material/styles';
2
+
3
+ let theme = createTheme({
4
+ palette: {
5
+ mode: 'dark',
6
+ primary: {
7
+ main: '#35C2D4',
8
+ light: '#73D7E3',
9
+ dark: '#1B98A8',
10
+ contrastText: '#061017',
11
+ },
12
+ secondary: {
13
+ main: '#9AA7BA',
14
+ light: '#C4CEDB',
15
+ dark: '#6E7C92',
16
+ contrastText: '#09101A',
17
+ },
18
+ background: {
19
+ default: '#090C12',
20
+ paper: '#121926',
21
+ },
22
+ text: {
23
+ primary: '#E8EDF5',
24
+ secondary: '#9DA9BC',
25
+ },
26
+ divider: 'rgba(194, 207, 228, 0.16)',
27
+ error: {
28
+ main: '#E36C61',
29
+ },
30
+ warning: {
31
+ main: '#E3A34B',
32
+ },
33
+ success: {
34
+ main: '#53C18A',
35
+ },
36
+ info: {
37
+ main: '#35C2D4',
38
+ },
39
+ },
40
+ shape: {
41
+ borderRadius: 12,
42
+ },
43
+ typography: {
44
+ fontFamily: [
45
+ 'Helvetica Neue',
46
+ 'Helvetica',
47
+ 'Arial',
48
+ 'sans-serif'
49
+ ].join(','),
50
+ h1: {
51
+ fontFamily: 'Helvetica Neue, Helvetica, Arial, sans-serif',
52
+ fontWeight: 1500,
53
+ },
54
+ h2: {
55
+ fontFamily: 'Helvetica Neue, Helvetica, Arial, sans-serif',
56
+ fontWeight: 200,
57
+ },
58
+ h3: {
59
+ fontFamily: 'Helvetica Neue, Helvetica, Arial, sans-serif',
60
+ fontWeight: 250,
61
+ },
62
+ h4: {
63
+ fontFamily: 'Helvetica Neue, Helvetica, Arial, sans-serif',
64
+ fontWeight: 300,
65
+ letterSpacing: '0.01em',
66
+ },
67
+ h5: {
68
+ fontFamily: 'Helvetica Neue, Helvetica, Arial, sans-serif',
69
+ fontWeight: 350,
70
+ letterSpacing: '0.01em',
71
+ },
72
+ h6: {
73
+ fontFamily: 'Helvetica Neue, Helvetica, Arial, sans-serif',
74
+ fontWeight: 400,
75
+ },
76
+ body1: {
77
+ fontFamily: 'Helvetica Neue, Helvetica, Arial, sans-serif',
78
+ fontWeight: 300,
79
+ },
80
+ body2: {
81
+ fontFamily: 'Helvetica Neue, Helvetica, Arial, sans-serif',
82
+ fontWeight: 300,
83
+ },
84
+ button: {
85
+ fontFamily: 'Helvetica Neue, Helvetica, Arial, sans-serif',
86
+ fontWeight: 400,
87
+ letterSpacing: '0.01em',
88
+ },
89
+ },
90
+ components: {
91
+ MuiCssBaseline: {
92
+ styleOverrides: {
93
+ ':root': {
94
+ colorScheme: 'dark',
95
+ },
96
+ body: {
97
+ margin: 0,
98
+ minHeight: '100vh',
99
+ backgroundColor: '#090C12',
100
+ backgroundImage: 'radial-gradient(1400px 700px at 8% -10%, rgba(53, 194, 212, 0.16), transparent 55%), radial-gradient(900px 500px at 92% -20%, rgba(83, 193, 138, 0.12), transparent 60%), linear-gradient(160deg, #090C12 0%, #0D121B 45%, #0A0F16 100%)',
101
+ color: '#E8EDF5',
102
+ },
103
+ '#root': {
104
+ minHeight: '100vh',
105
+ },
106
+ '*::-webkit-scrollbar': {
107
+ width: '10px',
108
+ height: '10px',
109
+ },
110
+ '*::-webkit-scrollbar-track': {
111
+ background: 'rgba(157, 169, 188, 0.14)',
112
+ borderRadius: '999px',
113
+ },
114
+ '*::-webkit-scrollbar-thumb': {
115
+ background: 'rgba(157, 169, 188, 0.45)',
116
+ borderRadius: '999px',
117
+ border: '2px solid rgba(0, 0, 0, 0)',
118
+ backgroundClip: 'padding-box',
119
+ '&:hover': {
120
+ background: 'rgba(157, 169, 188, 0.62)',
121
+ },
122
+ },
123
+ '*::-webkit-scrollbar-corner': {
124
+ background: 'rgba(157, 169, 188, 0.12)',
125
+ },
126
+ '*': {
127
+ scrollbarWidth: 'thin',
128
+ scrollbarColor: 'rgba(157, 169, 188, 0.45) rgba(157, 169, 188, 0.14)',
129
+ },
130
+ },
131
+ },
132
+ MuiPaper: {
133
+ styleOverrides: {
134
+ root: {
135
+ backgroundColor: '#121926',
136
+ backgroundImage: 'linear-gradient(180deg, rgba(20, 27, 40, 0.92) 0%, rgba(15, 22, 34, 0.95) 100%)',
137
+ border: '1px solid rgba(194, 207, 228, 0.14)',
138
+ boxShadow: '0 18px 32px rgba(4, 8, 14, 0.42)',
139
+ backdropFilter: 'blur(8px)',
140
+ },
141
+ },
142
+ },
143
+ MuiCard: {
144
+ styleOverrides: {
145
+ root: {
146
+ backgroundColor: '#111826',
147
+ backgroundImage: 'linear-gradient(180deg, rgba(18, 25, 38, 0.98) 0%, rgba(15, 22, 34, 0.96) 100%)',
148
+ border: '1px solid rgba(194, 207, 228, 0.14)',
149
+ boxShadow: '0 10px 18px rgba(4, 8, 14, 0.3)',
150
+ transition: 'border-color 180ms ease, box-shadow 180ms ease, transform 180ms ease',
151
+ '&:hover': {
152
+ borderColor: 'rgba(115, 215, 227, 0.3)',
153
+ boxShadow: '0 16px 30px rgba(4, 8, 14, 0.46)',
154
+ transform: 'translateY(-1px)',
155
+ },
156
+ },
157
+ },
158
+ },
159
+ MuiButton: {
160
+ styleOverrides: {
161
+ root: {
162
+ textTransform: 'none',
163
+ borderRadius: 10,
164
+ fontWeight: 600,
165
+ paddingInline: 16,
166
+ display: 'inline-flex',
167
+ alignItems: 'center',
168
+ justifyContent: 'center',
169
+ lineHeight: 1.2,
170
+ '& .MuiButton-startIcon, & .MuiButton-endIcon': {
171
+ display: 'inline-flex',
172
+ alignItems: 'center',
173
+ },
174
+ transition: 'transform 160ms ease, box-shadow 160ms ease, border-color 160ms ease, background-color 160ms ease',
175
+ },
176
+ contained: {
177
+ boxShadow: '0 8px 18px rgba(6, 10, 18, 0.46)',
178
+ '&:hover': {
179
+ boxShadow: '0 12px 22px rgba(6, 10, 18, 0.58)',
180
+ transform: 'translateY(-1px)',
181
+ },
182
+ },
183
+ containedPrimary: {
184
+ backgroundImage: 'linear-gradient(135deg, #35C2D4 0%, #2AA9B9 55%, #228E9D 100%)',
185
+ },
186
+ containedError: {
187
+ backgroundImage: 'linear-gradient(135deg, #E36C61 0%, #CF5A4E 100%)',
188
+ },
189
+ outlined: {
190
+ borderColor: 'rgba(157, 169, 188, 0.4)',
191
+ '&:hover': {
192
+ borderColor: '#35C2D4',
193
+ backgroundColor: 'rgba(53, 194, 212, 0.08)',
194
+ },
195
+ },
196
+ },
197
+ },
198
+ MuiInputBase: {
199
+ styleOverrides: {
200
+ root: {
201
+ '&:not(.MuiInputBase-multiline)': {
202
+ alignItems: 'center',
203
+ },
204
+ },
205
+ input: {
206
+ lineHeight: 1.4,
207
+ },
208
+ },
209
+ },
210
+ MuiTextField: {
211
+ styleOverrides: {
212
+ root: {
213
+ '& .MuiOutlinedInput-root': {
214
+ backgroundColor: 'rgba(10, 15, 23, 0.84)',
215
+ '& fieldset': {
216
+ borderColor: 'rgba(157, 169, 188, 0.3)',
217
+ },
218
+ '&:hover fieldset': {
219
+ borderColor: 'rgba(157, 169, 188, 0.55)',
220
+ },
221
+ '&.Mui-focused fieldset': {
222
+ borderColor: '#35C2D4',
223
+ },
224
+ },
225
+ },
226
+ },
227
+ },
228
+ MuiSelect: {
229
+ styleOverrides: {
230
+ root: {
231
+ backgroundColor: 'rgba(10, 15, 23, 0.84)',
232
+ '& .MuiOutlinedInput-notchedOutline': {
233
+ borderColor: 'rgba(157, 169, 188, 0.3)',
234
+ },
235
+ '&:hover .MuiOutlinedInput-notchedOutline': {
236
+ borderColor: 'rgba(157, 169, 188, 0.55)',
237
+ },
238
+ '&.Mui-focused .MuiOutlinedInput-notchedOutline': {
239
+ borderColor: '#35C2D4',
240
+ },
241
+ },
242
+ select: {
243
+ display: 'flex',
244
+ alignItems: 'center',
245
+ },
246
+ },
247
+ },
248
+ MuiMenuItem: {
249
+ styleOverrides: {
250
+ root: {
251
+ backgroundColor: '#121926',
252
+ '&:hover': {
253
+ backgroundColor: 'rgba(53, 194, 212, 0.08)',
254
+ },
255
+ '&.Mui-selected': {
256
+ backgroundColor: 'rgba(53, 194, 212, 0.14)',
257
+ '&:hover': {
258
+ backgroundColor: 'rgba(53, 194, 212, 0.2)',
259
+ },
260
+ },
261
+ },
262
+ },
263
+ },
264
+ MuiChip: {
265
+ styleOverrides: {
266
+ root: {
267
+ backgroundColor: 'rgba(157, 169, 188, 0.16)',
268
+ color: '#E8EDF5',
269
+ border: '1px solid rgba(157, 169, 188, 0.24)',
270
+ '&.MuiChip-colorPrimary': {
271
+ backgroundColor: 'rgba(53, 194, 212, 0.2)',
272
+ color: '#C8F3F9',
273
+ },
274
+ },
275
+ outlined: {
276
+ borderColor: 'rgba(157, 169, 188, 0.4)',
277
+ '&.MuiChip-colorPrimary': {
278
+ borderColor: '#35C2D4',
279
+ color: '#73D7E3',
280
+ },
281
+ },
282
+ },
283
+ },
284
+ MuiAccordion: {
285
+ styleOverrides: {
286
+ root: {
287
+ backgroundColor: 'rgba(12, 18, 28, 0.7)',
288
+ border: '1px solid rgba(194, 207, 228, 0.16)',
289
+ borderRadius: 12,
290
+ overflow: 'hidden',
291
+ '&:before': {
292
+ display: 'none',
293
+ },
294
+ '&.Mui-expanded': {
295
+ margin: 0,
296
+ },
297
+ },
298
+ },
299
+ },
300
+ MuiAccordionSummary: {
301
+ styleOverrides: {
302
+ root: {
303
+ backgroundColor: 'rgba(15, 22, 34, 0.8)',
304
+ borderRadius: 12,
305
+ minHeight: 44,
306
+ '& .MuiAccordionSummary-content': {
307
+ margin: '10px 0',
308
+ alignItems: 'center',
309
+ },
310
+ '&.Mui-expanded': {
311
+ minHeight: 44,
312
+ },
313
+ '&.Mui-expanded .MuiAccordionSummary-content': {
314
+ margin: '10px 0',
315
+ },
316
+ '&:hover': {
317
+ backgroundColor: 'rgba(19, 28, 42, 0.9)',
318
+ },
319
+ },
320
+ },
321
+ },
322
+ MuiDialog: {
323
+ styleOverrides: {
324
+ paper: {
325
+ backgroundColor: '#121926',
326
+ backgroundImage: 'linear-gradient(180deg, rgba(20, 27, 40, 0.98) 0%, rgba(14, 21, 33, 0.98) 100%)',
327
+ border: '1px solid rgba(194, 207, 228, 0.18)',
328
+ borderRadius: 14,
329
+ boxShadow: '0 28px 48px rgba(4, 8, 14, 0.6)',
330
+ },
331
+ },
332
+ },
333
+ MuiDialogTitle: {
334
+ styleOverrides: {
335
+ root: {
336
+ backgroundColor: 'rgba(14, 21, 33, 0.8)',
337
+ borderBottom: '1px solid rgba(194, 207, 228, 0.15)',
338
+ color: '#F4F7FC',
339
+ fontWeight: 600,
340
+ fontSize: '1.15rem',
341
+ },
342
+ },
343
+ },
344
+ MuiDialogContent: {
345
+ styleOverrides: {
346
+ root: {
347
+ backgroundColor: 'rgba(14, 21, 33, 0.64)',
348
+ color: '#CCD5E3',
349
+ },
350
+ },
351
+ },
352
+ MuiDialogActions: {
353
+ styleOverrides: {
354
+ root: {
355
+ backgroundColor: 'rgba(14, 21, 33, 0.72)',
356
+ borderTop: '1px solid rgba(194, 207, 228, 0.15)',
357
+ padding: '14px 20px',
358
+ gap: 8,
359
+ },
360
+ },
361
+ },
362
+ MuiListItem: {
363
+ styleOverrides: {
364
+ root: {
365
+ '&:hover': {
366
+ backgroundColor: 'rgba(53, 194, 212, 0.08)',
367
+ },
368
+ '&.Mui-selected': {
369
+ backgroundColor: 'rgba(53, 194, 212, 0.14)',
370
+ '&:hover': {
371
+ backgroundColor: 'rgba(53, 194, 212, 0.2)',
372
+ },
373
+ },
374
+ },
375
+ },
376
+ },
377
+ MuiCheckbox: {
378
+ styleOverrides: {
379
+ root: {
380
+ color: '#9AA7BA',
381
+ '&.Mui-checked': {
382
+ color: '#35C2D4',
383
+ },
384
+ '&:hover': {
385
+ backgroundColor: 'rgba(53, 194, 212, 0.08)',
386
+ },
387
+ },
388
+ },
389
+ },
390
+ MuiFormControlLabel: {
391
+ styleOverrides: {
392
+ label: {
393
+ color: '#CCD5E3',
394
+ fontSize: '0.875rem',
395
+ },
396
+ },
397
+ },
398
+ MuiSlider: {
399
+ styleOverrides: {
400
+ root: {
401
+ color: '#35C2D4',
402
+ },
403
+ rail: {
404
+ backgroundColor: 'rgba(157, 169, 188, 0.24)',
405
+ },
406
+ track: {
407
+ backgroundColor: '#35C2D4',
408
+ border: 0,
409
+ },
410
+ thumb: {
411
+ backgroundColor: '#74DEE9',
412
+ '&:hover': {
413
+ boxShadow: '0 0 0 8px rgba(53, 194, 212, 0.2)',
414
+ },
415
+ },
416
+ },
417
+ },
418
+ MuiLinearProgress: {
419
+ styleOverrides: {
420
+ root: {
421
+ backgroundColor: 'rgba(157, 169, 188, 0.2)',
422
+ },
423
+ bar: {
424
+ backgroundColor: '#35C2D4',
425
+ },
426
+ },
427
+ },
428
+ MuiCircularProgress: {
429
+ styleOverrides: {
430
+ root: {
431
+ color: '#35C2D4',
432
+ },
433
+ },
434
+ },
435
+ MuiTabs: {
436
+ styleOverrides: {
437
+ root: {
438
+ '& .MuiTabs-indicator': {
439
+ backgroundColor: '#35C2D4',
440
+ },
441
+ },
442
+ },
443
+ },
444
+ MuiTab: {
445
+ styleOverrides: {
446
+ root: {
447
+ color: '#9DA9BC',
448
+ '&.Mui-selected': {
449
+ color: '#35C2D4',
450
+ },
451
+ '&:hover': {
452
+ color: '#E8EDF5',
453
+ },
454
+ },
455
+ },
456
+ },
457
+ MuiBackdrop: {
458
+ styleOverrides: {
459
+ root: {
460
+ backgroundColor: 'rgba(5, 9, 16, 0.84)',
461
+ backdropFilter: 'blur(4px)',
462
+ },
463
+ },
464
+ },
465
+ MuiDivider: {
466
+ styleOverrides: {
467
+ root: {
468
+ borderColor: 'rgba(194, 207, 228, 0.16)',
469
+ },
470
+ },
471
+ },
472
+ MuiIconButton: {
473
+ styleOverrides: {
474
+ root: {
475
+ color: '#9DA9BC',
476
+ '&:hover': {
477
+ backgroundColor: 'rgba(53, 194, 212, 0.1)',
478
+ color: '#73D7E3',
479
+ },
480
+ },
481
+ },
482
+ },
483
+ MuiContainer: {
484
+ styleOverrides: {
485
+ root: {
486
+ backgroundColor: 'transparent',
487
+ background: 'transparent',
488
+ },
489
+ },
490
+ },
491
+ },
492
+ });
493
+
494
+ theme = responsiveFontSizes(theme, {
495
+ breakpoints: ['sm', 'md', 'lg'],
496
+ factor: 2.4,
497
+ });
498
+
499
+ export const lightTheme = createTheme(theme, {
500
+ palette: {
501
+ mode: 'light',
502
+ primary: {
503
+ main: '#1497A8',
504
+ light: '#4CBCCA',
505
+ dark: '#0F7482',
506
+ contrastText: '#F7FDFF',
507
+ },
508
+ secondary: {
509
+ main: '#64748B',
510
+ light: '#93A3B8',
511
+ dark: '#475569',
512
+ contrastText: '#F8FAFC',
513
+ },
514
+ background: {
515
+ default: '#F5F9FC',
516
+ paper: '#FFFFFF',
517
+ },
518
+ text: {
519
+ primary: '#0F172A',
520
+ secondary: '#475569',
521
+ },
522
+ divider: 'rgba(15, 23, 42, 0.14)',
523
+ error: {
524
+ main: '#DC5B57',
525
+ },
526
+ warning: {
527
+ main: '#D08C30',
528
+ },
529
+ success: {
530
+ main: '#2E9E63',
531
+ },
532
+ info: {
533
+ main: '#1497A8',
534
+ },
535
+ },
536
+ components: {
537
+ MuiCssBaseline: {
538
+ styleOverrides: {
539
+ ':root': {
540
+ colorScheme: 'light',
541
+ },
542
+ body: {
543
+ margin: 0,
544
+ minHeight: '100vh',
545
+ backgroundColor: '#F5F9FC',
546
+ backgroundImage: 'radial-gradient(1400px 700px at 8% -10%, rgba(20, 151, 168, 0.14), transparent 55%), radial-gradient(900px 500px at 92% -20%, rgba(46, 158, 99, 0.1), transparent 60%), linear-gradient(160deg, #F6FAFD 0%, #EFF5FA 45%, #F8FBFE 100%)',
547
+ color: '#0F172A',
548
+ },
549
+ '#root': {
550
+ minHeight: '100vh',
551
+ },
552
+ '*::-webkit-scrollbar-track': {
553
+ background: 'rgba(100, 116, 139, 0.12)',
554
+ borderRadius: '999px',
555
+ },
556
+ '*::-webkit-scrollbar-thumb': {
557
+ background: 'rgba(100, 116, 139, 0.38)',
558
+ borderRadius: '999px',
559
+ '&:hover': {
560
+ background: 'rgba(100, 116, 139, 0.52)',
561
+ },
562
+ },
563
+ '*::-webkit-scrollbar-corner': {
564
+ background: 'rgba(100, 116, 139, 0.1)',
565
+ },
566
+ '*': {
567
+ scrollbarWidth: 'thin',
568
+ scrollbarColor: 'rgba(100, 116, 139, 0.38) rgba(100, 116, 139, 0.12)',
569
+ },
570
+ },
571
+ },
572
+ MuiPaper: {
573
+ styleOverrides: {
574
+ root: {
575
+ backgroundColor: '#FFFFFF',
576
+ backgroundImage: 'linear-gradient(180deg, rgba(255, 255, 255, 0.98) 0%, rgba(248, 251, 255, 0.98) 100%)',
577
+ border: '1px solid rgba(15, 23, 42, 0.12)',
578
+ boxShadow: '0 14px 26px rgba(15, 23, 42, 0.08)',
579
+ },
580
+ },
581
+ },
582
+ MuiCard: {
583
+ styleOverrides: {
584
+ root: {
585
+ backgroundColor: '#FFFFFF',
586
+ backgroundImage: 'linear-gradient(180deg, rgba(255, 255, 255, 0.99) 0%, rgba(248, 251, 255, 0.99) 100%)',
587
+ border: '1px solid rgba(15, 23, 42, 0.12)',
588
+ boxShadow: '0 10px 18px rgba(15, 23, 42, 0.08)',
589
+ '&:hover': {
590
+ borderColor: 'rgba(20, 151, 168, 0.32)',
591
+ boxShadow: '0 16px 30px rgba(15, 23, 42, 0.12)',
592
+ },
593
+ },
594
+ },
595
+ },
596
+ MuiButton: {
597
+ styleOverrides: {
598
+ contained: {
599
+ boxShadow: '0 8px 18px rgba(20, 151, 168, 0.22)',
600
+ '&:hover': {
601
+ boxShadow: '0 12px 24px rgba(20, 151, 168, 0.28)',
602
+ },
603
+ },
604
+ containedPrimary: {
605
+ backgroundImage: 'linear-gradient(135deg, #1497A8 0%, #1AAABC 55%, #107A88 100%)',
606
+ },
607
+ containedError: {
608
+ backgroundImage: 'linear-gradient(135deg, #DC5B57 0%, #CB4B45 100%)',
609
+ },
610
+ outlined: {
611
+ borderColor: 'rgba(100, 116, 139, 0.32)',
612
+ '&:hover': {
613
+ borderColor: '#1497A8',
614
+ backgroundColor: 'rgba(20, 151, 168, 0.08)',
615
+ },
616
+ },
617
+ },
618
+ },
619
+ MuiTextField: {
620
+ styleOverrides: {
621
+ root: {
622
+ '& .MuiOutlinedInput-root': {
623
+ backgroundColor: 'rgba(255, 255, 255, 0.92)',
624
+ '& fieldset': {
625
+ borderColor: 'rgba(100, 116, 139, 0.28)',
626
+ },
627
+ '&:hover fieldset': {
628
+ borderColor: 'rgba(100, 116, 139, 0.5)',
629
+ },
630
+ '&.Mui-focused fieldset': {
631
+ borderColor: '#1497A8',
632
+ },
633
+ },
634
+ },
635
+ },
636
+ },
637
+ MuiSelect: {
638
+ styleOverrides: {
639
+ root: {
640
+ backgroundColor: 'rgba(255, 255, 255, 0.92)',
641
+ '& .MuiOutlinedInput-notchedOutline': {
642
+ borderColor: 'rgba(100, 116, 139, 0.28)',
643
+ },
644
+ '&:hover .MuiOutlinedInput-notchedOutline': {
645
+ borderColor: 'rgba(100, 116, 139, 0.5)',
646
+ },
647
+ '&.Mui-focused .MuiOutlinedInput-notchedOutline': {
648
+ borderColor: '#1497A8',
649
+ },
650
+ },
651
+ },
652
+ },
653
+ MuiMenuItem: {
654
+ styleOverrides: {
655
+ root: {
656
+ backgroundColor: '#FFFFFF',
657
+ '&:hover': {
658
+ backgroundColor: 'rgba(20, 151, 168, 0.08)',
659
+ },
660
+ '&.Mui-selected': {
661
+ backgroundColor: 'rgba(20, 151, 168, 0.14)',
662
+ '&:hover': {
663
+ backgroundColor: 'rgba(20, 151, 168, 0.2)',
664
+ },
665
+ },
666
+ },
667
+ },
668
+ },
669
+ MuiChip: {
670
+ styleOverrides: {
671
+ root: {
672
+ backgroundColor: 'rgba(100, 116, 139, 0.12)',
673
+ color: '#0F172A',
674
+ border: '1px solid rgba(100, 116, 139, 0.24)',
675
+ },
676
+ },
677
+ },
678
+ MuiAccordion: {
679
+ styleOverrides: {
680
+ root: {
681
+ backgroundColor: 'rgba(255, 255, 255, 0.82)',
682
+ border: '1px solid rgba(15, 23, 42, 0.12)',
683
+ borderRadius: 12,
684
+ overflow: 'hidden',
685
+ '&:before': {
686
+ display: 'none',
687
+ },
688
+ '&.Mui-expanded': {
689
+ margin: 0,
690
+ },
691
+ },
692
+ },
693
+ },
694
+ MuiAccordionSummary: {
695
+ styleOverrides: {
696
+ root: {
697
+ backgroundColor: 'rgba(248, 251, 255, 0.95)',
698
+ borderRadius: 12,
699
+ '&:hover': {
700
+ backgroundColor: 'rgba(239, 245, 250, 1)',
701
+ },
702
+ },
703
+ },
704
+ },
705
+ MuiDialog: {
706
+ styleOverrides: {
707
+ paper: {
708
+ backgroundColor: '#FFFFFF',
709
+ backgroundImage: 'linear-gradient(180deg, rgba(255, 255, 255, 0.99) 0%, rgba(246, 250, 254, 0.99) 100%)',
710
+ border: '1px solid rgba(15, 23, 42, 0.12)',
711
+ boxShadow: '0 28px 48px rgba(15, 23, 42, 0.16)',
712
+ },
713
+ },
714
+ },
715
+ MuiDialogTitle: {
716
+ styleOverrides: {
717
+ root: {
718
+ backgroundColor: 'rgba(246, 250, 254, 0.98)',
719
+ borderBottom: '1px solid rgba(15, 23, 42, 0.1)',
720
+ color: '#0F172A',
721
+ },
722
+ },
723
+ },
724
+ MuiDialogContent: {
725
+ styleOverrides: {
726
+ root: {
727
+ backgroundColor: 'rgba(255, 255, 255, 0.98)',
728
+ color: '#334155',
729
+ },
730
+ },
731
+ },
732
+ MuiDialogActions: {
733
+ styleOverrides: {
734
+ root: {
735
+ backgroundColor: 'rgba(246, 250, 254, 0.98)',
736
+ borderTop: '1px solid rgba(15, 23, 42, 0.1)',
737
+ },
738
+ },
739
+ },
740
+ MuiListItem: {
741
+ styleOverrides: {
742
+ root: {
743
+ '&:hover': {
744
+ backgroundColor: 'rgba(53, 194, 212, 0.08)',
745
+ },
746
+ '&.Mui-selected': {
747
+ backgroundColor: 'rgba(53, 194, 212, 0.14)',
748
+ '&:hover': {
749
+ backgroundColor: 'rgba(53, 194, 212, 0.2)',
750
+ },
751
+ },
752
+ },
753
+ },
754
+ },
755
+ MuiCheckbox: {
756
+ styleOverrides: {
757
+ root: {
758
+ color: '#64748B',
759
+ '&.Mui-checked': {
760
+ color: '#1497A8',
761
+ },
762
+ '&:hover': {
763
+ backgroundColor: 'rgba(20, 151, 168, 0.08)',
764
+ },
765
+ },
766
+ },
767
+ },
768
+ MuiFormControlLabel: {
769
+ styleOverrides: {
770
+ label: {
771
+ color: '#334155',
772
+ },
773
+ },
774
+ },
775
+ MuiSlider: {
776
+ styleOverrides: {
777
+ root: {
778
+ color: '#1497A8',
779
+ },
780
+ rail: {
781
+ backgroundColor: 'rgba(100, 116, 139, 0.24)',
782
+ },
783
+ track: {
784
+ backgroundColor: '#1497A8',
785
+ border: 0,
786
+ },
787
+ thumb: {
788
+ backgroundColor: '#4CBCCA',
789
+ '&:hover': {
790
+ boxShadow: '0 0 0 8px rgba(20, 151, 168, 0.2)',
791
+ },
792
+ },
793
+ },
794
+ },
795
+ MuiLinearProgress: {
796
+ styleOverrides: {
797
+ root: {
798
+ backgroundColor: 'rgba(100, 116, 139, 0.2)',
799
+ },
800
+ bar: {
801
+ backgroundColor: '#1497A8',
802
+ },
803
+ },
804
+ },
805
+ MuiCircularProgress: {
806
+ styleOverrides: {
807
+ root: {
808
+ color: '#1497A8',
809
+ },
810
+ },
811
+ },
812
+ MuiTabs: {
813
+ styleOverrides: {
814
+ root: {
815
+ '& .MuiTabs-indicator': {
816
+ backgroundColor: '#1497A8',
817
+ },
818
+ },
819
+ },
820
+ },
821
+ MuiTab: {
822
+ styleOverrides: {
823
+ root: {
824
+ color: '#64748B',
825
+ '&.Mui-selected': {
826
+ color: '#1497A8',
827
+ },
828
+ '&:hover': {
829
+ color: '#0F172A',
830
+ },
831
+ },
832
+ },
833
+ },
834
+ MuiBackdrop: {
835
+ styleOverrides: {
836
+ root: {
837
+ backgroundColor: 'rgba(15, 23, 42, 0.38)',
838
+ backdropFilter: 'blur(4px)',
839
+ },
840
+ },
841
+ },
842
+ MuiDivider: {
843
+ styleOverrides: {
844
+ root: {
845
+ borderColor: 'rgba(15, 23, 42, 0.14)',
846
+ },
847
+ },
848
+ },
849
+ MuiIconButton: {
850
+ styleOverrides: {
851
+ root: {
852
+ color: '#64748B',
853
+ '&:hover': {
854
+ backgroundColor: 'rgba(20, 151, 168, 0.1)',
855
+ color: '#1497A8',
856
+ },
857
+ },
858
+ },
859
+ },
860
+ },
861
+ });
862
+
863
+ export const appStyles = {
864
+ root: {
865
+ minHeight: '100vh',
866
+ background: 'transparent',
867
+ backgroundColor: 'transparent',
868
+ overflow: 'visible',
869
+ position: 'relative',
870
+ display: 'flex',
871
+ flexDirection: 'column',
872
+ },
873
+ container: (showWelcomePage) => ({
874
+ py: { xs: 1, sm: 1.5, md: 2.5 },
875
+ px: { xs: 0.75, sm: 1.25, md: 2.5 },
876
+ minHeight: '100vh',
877
+ display: 'flex',
878
+ flexDirection: 'column',
879
+ backgroundColor: 'transparent',
880
+ background: 'transparent',
881
+ overflow: 'visible',
882
+ boxSizing: 'border-box',
883
+ width: '100%',
884
+ maxWidth: '100%',
885
+ filter: showWelcomePage ? 'blur(8px)' : 'none',
886
+ transition: 'filter 0.3s ease-in-out',
887
+ }),
888
+ headerRow: {
889
+ display: 'flex',
890
+ justifyContent: 'space-between',
891
+ alignItems: { xs: 'stretch', md: 'flex-start' },
892
+ flexDirection: { xs: 'column', md: 'row' },
893
+ gap: { xs: 1.25, sm: 1.75, md: 2 },
894
+ mb: { xs: 1, sm: 1.5 },
895
+ },
896
+ headerBrand: {
897
+ position: 'relative',
898
+ display: 'flex',
899
+ alignItems: 'center',
900
+ gap: { xs: 1.25, sm: 2 },
901
+ py: { xs: 0.25, sm: 0.5 },
902
+ },
903
+ logo: {
904
+ width: 60,
905
+ height: 60,
906
+ backgroundImage: 'url(/fragmenta_icon_1024.png)',
907
+ backgroundSize: 'cover',
908
+ backgroundPosition: 'center',
909
+ borderRadius: 2,
910
+ border: '1px solid rgba(194, 207, 228, 0.22)',
911
+ boxShadow: '0 10px 20px rgba(4, 8, 14, 0.36)',
912
+ filter: 'drop-shadow(0 4px 8px rgba(0, 0, 0, 0.3))',
913
+ },
914
+ title: {
915
+ color: 'text.primary',
916
+ fontFamily: '"Bitcount Single", "IBM Plex Mono", "JetBrains Mono", "Space Mono", "Courier New", monospace',
917
+ fontWeight: 400,
918
+ letterSpacing: '0.02em',
919
+ textShadow: '0 2px 10px rgba(0, 0, 0, 0.6)',
920
+ },
921
+ headerActionsContainer: (isCompactLayout) => ({
922
+ display: 'flex',
923
+ alignItems: 'stretch',
924
+ justifyContent: { xs: 'flex-start', md: 'flex-end' },
925
+ gap: { xs: 1, sm: 1.5 },
926
+ flexDirection: isCompactLayout ? 'column' : 'row',
927
+ flexWrap: isCompactLayout ? 'wrap' : 'nowrap',
928
+ width: { xs: '100%', md: 'auto' },
929
+ }),
930
+ headerActionsGrid: (isCompactLayout) => ({
931
+ display: 'grid',
932
+ gridTemplateColumns: isCompactLayout
933
+ ? 'repeat(2, minmax(0, 1fr))'
934
+ : 'repeat(2, 122px)',
935
+ gap: { xs: 0.75, sm: 1 },
936
+ justifyContent: 'flex-end',
937
+ flex: isCompactLayout ? '1 1 auto' : '0 1 auto',
938
+ width: isCompactLayout ? '100%' : 'auto',
939
+ }),
940
+ headerActionButton: {
941
+ fontSize: { xs: '0.70rem', sm: '0.72rem' },
942
+ height: { xs: 34, sm: 36 },
943
+ minWidth: 0,
944
+ width: '100%',
945
+ px: { xs: 1, sm: 1.5 },
946
+ '& .MuiButton-startIcon svg': {
947
+ width: { xs: 14, sm: 15 },
948
+ height: { xs: 14, sm: 15 },
949
+ },
950
+ },
951
+ gpuCard: (isCompactLayout) => ({
952
+ p: { xs: 1.25, sm: 1.75 },
953
+ bgcolor: 'background.paper',
954
+ borderRadius: 2.5,
955
+ border: '1px solid',
956
+ borderColor: 'divider',
957
+ minWidth: isCompactLayout ? '100%' : 270,
958
+ flexShrink: 0,
959
+ position: 'relative',
960
+ overflow: 'hidden',
961
+ boxShadow: '0 16px 32px rgba(4, 8, 14, 0.44)',
962
+ }),
963
+ gpuUsageTrack: {
964
+ position: 'relative',
965
+ width: '100%',
966
+ height: 6,
967
+ bgcolor: 'rgba(157, 169, 188, 0.2)',
968
+ borderRadius: 3,
969
+ overflow: 'hidden',
970
+ },
971
+ gpuUsageFill: (width, color) => ({
972
+ position: 'absolute',
973
+ top: 0,
974
+ left: 0,
975
+ height: '100%',
976
+ width,
977
+ bgcolor: color,
978
+ borderRadius: 3,
979
+ transition: 'width 0.3s ease-in-out',
980
+ }),
981
+ emphasizedPrimaryBody2: {
982
+ fontWeight: 'bold',
983
+ color: 'primary.main',
984
+ },
985
+ advancedSettingsDetails: (muiTheme) => {
986
+ const isDark = muiTheme.palette.mode === 'dark';
987
+ return {
988
+ backgroundColor: isDark
989
+ ? 'rgba(10, 15, 23, 0.46)'
990
+ : 'rgba(255, 255, 255, 0.82)',
991
+ borderTop: isDark
992
+ ? '1px solid rgba(194, 207, 228, 0.12)'
993
+ : '1px solid rgba(15, 23, 42, 0.08)',
994
+ borderBottomLeftRadius: 12,
995
+ borderBottomRightRadius: 12,
996
+ maxHeight: { xs: 'none', md: '400px' },
997
+ overflowY: { xs: 'visible', md: 'auto' },
998
+ overflowX: 'hidden',
999
+ '&::-webkit-scrollbar': {
1000
+ width: '8px',
1001
+ },
1002
+ '&::-webkit-scrollbar-track': {
1003
+ background: isDark
1004
+ ? 'rgba(157, 169, 188, 0.14)'
1005
+ : 'rgba(100, 116, 139, 0.14)',
1006
+ borderRadius: '4px',
1007
+ },
1008
+ '&::-webkit-scrollbar-thumb': {
1009
+ background: isDark
1010
+ ? 'rgba(157, 169, 188, 0.45)'
1011
+ : 'rgba(100, 116, 139, 0.42)',
1012
+ borderRadius: '4px',
1013
+ '&:hover': {
1014
+ background: isDark
1015
+ ? 'rgba(157, 169, 188, 0.62)'
1016
+ : 'rgba(100, 116, 139, 0.56)',
1017
+ },
1018
+ },
1019
+ };
1020
+ },
1021
+ mainLayout: {
1022
+ display: 'flex',
1023
+ flexDirection: { xs: 'column', md: 'row' },
1024
+ width: '100%',
1025
+ flex: 1,
1026
+ gap: { xs: 1, sm: 1.25, md: 1.5 },
1027
+ borderRadius: 3,
1028
+ minHeight: 0,
1029
+ },
1030
+ navPaper: {
1031
+ width: { xs: '100%', md: 220 },
1032
+ backgroundColor: 'background.paper',
1033
+ borderRadius: 2.5,
1034
+ overflow: 'hidden',
1035
+ display: 'flex',
1036
+ flexDirection: 'column',
1037
+ height: '100%',
1038
+ },
1039
+ navigationTabs: (isCompactLayout) => ({
1040
+ height: isCompactLayout ? 'auto' : '100%',
1041
+ p: { xs: 0.5, sm: 1 },
1042
+ gap: { xs: 0.25, sm: 0.5 },
1043
+ '& .MuiTabs-indicator': {
1044
+ display: 'none',
1045
+ },
1046
+ '& .MuiTab-root': {
1047
+ alignItems: 'center',
1048
+ justifyContent: isCompactLayout ? 'center' : 'flex-start',
1049
+ textAlign: isCompactLayout ? 'center' : 'left',
1050
+ minHeight: { xs: 40, sm: 46 },
1051
+ fontSize: { xs: '0.78rem', sm: '0.86rem' },
1052
+ fontWeight: 500,
1053
+ textTransform: 'none',
1054
+ color: 'text.secondary',
1055
+ borderRadius: 2,
1056
+ px: { xs: 1, sm: 1.5 },
1057
+ py: { xs: 0.75, sm: 1 },
1058
+ mx: isCompactLayout ? 0.25 : 0,
1059
+ '& .MuiTab-iconWrapper': {
1060
+ marginBottom: '0 !important',
1061
+ },
1062
+ '&.Mui-selected': {
1063
+ color: 'primary.main',
1064
+ fontWeight: 600,
1065
+ backgroundColor: 'rgba(53, 194, 212, 0.16)',
1066
+ boxShadow: 'inset 0 0 0 1px rgba(115, 215, 227, 0.35)',
1067
+ },
1068
+ '&:hover': {
1069
+ color: 'text.primary',
1070
+ backgroundColor: 'rgba(53, 194, 212, 0.08)',
1071
+ },
1072
+ },
1073
+ }),
1074
+ mainContentPaper: (muiTheme) => ({
1075
+ flex: 1,
1076
+ backgroundColor: 'background.paper',
1077
+ borderRadius: 2.5,
1078
+ display: 'flex',
1079
+ flexDirection: 'column',
1080
+ minHeight: { xs: 'auto', md: 0 },
1081
+ overflow: 'visible',
1082
+ '& .MuiTypography-body1, & .MuiTypography-body2, & .MuiTypography-caption, & .MuiTypography-subtitle2': {
1083
+ fontSize: muiTheme.typography.body2.fontSize,
1084
+ lineHeight: muiTheme.typography.body2.lineHeight,
1085
+ },
1086
+ '& .MuiButton-startIcon svg, & .MuiButton-endIcon svg, & .MuiIconButton-root svg, & .MuiAccordionSummary-expandIconWrapper svg': {
1087
+ width: 20,
1088
+ height: 20,
1089
+ },
1090
+ }),
1091
+ elevatedInfoCard: (muiTheme) => {
1092
+ const isDark = muiTheme.palette.mode === 'dark';
1093
+ return {
1094
+ p: { xs: 1.5, sm: 2 },
1095
+ mb: 2,
1096
+ boxShadow: isDark
1097
+ ? '0 14px 28px rgba(4, 8, 14, 0.44)'
1098
+ : '0 14px 26px rgba(15, 23, 42, 0.1)',
1099
+ borderRadius: 2.5,
1100
+ border: isDark
1101
+ ? '1px solid rgba(194, 207, 228, 0.16)'
1102
+ : '1px solid rgba(15, 23, 42, 0.12)',
1103
+ background: isDark
1104
+ ? 'linear-gradient(160deg, rgba(17, 24, 37, 0.96) 0%, rgba(13, 20, 31, 0.92) 100%)'
1105
+ : 'linear-gradient(160deg, rgba(255, 255, 255, 0.98) 0%, rgba(245, 250, 255, 0.98) 100%)',
1106
+ '&:hover': {
1107
+ boxShadow: isDark
1108
+ ? '0 20px 34px rgba(4, 8, 14, 0.56)'
1109
+ : '0 20px 34px rgba(15, 23, 42, 0.14)',
1110
+ transform: 'translateY(-1px)',
1111
+ transition: 'all 0.3s ease',
1112
+ },
1113
+ transition: 'all 0.3s ease',
1114
+ };
1115
+ },
1116
+ modelMissingAlert: {
1117
+ mt: 2,
1118
+ backgroundColor: 'rgba(219, 80, 68, 0)',
1119
+ border: '1px solid #DB5044',
1120
+ borderRadius: 2,
1121
+ '& .MuiAlert-icon': {
1122
+ color: '#DB5044',
1123
+ },
1124
+ },
1125
+ modelMissingLink: {
1126
+ fontWeight: 600,
1127
+ color: 'primary.main',
1128
+ textDecoration: 'none',
1129
+ '&:hover': {
1130
+ textDecoration: 'underline',
1131
+ },
1132
+ },
1133
+ headerActionButtonWithOpacity: (isEnabled) => ({
1134
+ fontSize: { xs: '0.70rem', sm: '0.72rem' },
1135
+ height: { xs: 34, sm: 36 },
1136
+ minWidth: 0,
1137
+ width: '100%',
1138
+ px: { xs: 1, sm: 1.5 },
1139
+ opacity: isEnabled ? 1 : 0.5,
1140
+ '& .MuiButton-startIcon svg': {
1141
+ width: { xs: 14, sm: 15 },
1142
+ height: { xs: 14, sm: 15 },
1143
+ },
1144
+ }),
1145
+ gpuHeaderRow: {
1146
+ display: 'flex',
1147
+ alignItems: 'center',
1148
+ justifyContent: 'space-between',
1149
+ mb: 1,
1150
+ },
1151
+ gpuLabel: {
1152
+ fontWeight: 500,
1153
+ },
1154
+ gpuStatusGroup: {
1155
+ display: 'flex',
1156
+ alignItems: 'center',
1157
+ gap: 0.5,
1158
+ },
1159
+ gpuStatusDot: (status, animate = true) => ({
1160
+ width: 6,
1161
+ height: 6,
1162
+ borderRadius: '50%',
1163
+ bgcolor: status === 'good'
1164
+ ? 'success.main'
1165
+ : status === 'low'
1166
+ ? 'warning.main'
1167
+ : 'error.main',
1168
+ animation: animate ? 'pulse 2s infinite' : 'none',
1169
+ '@keyframes pulse': {
1170
+ '0%': { opacity: 1 },
1171
+ '50%': { opacity: 0.5 },
1172
+ '100%': { opacity: 1 },
1173
+ },
1174
+ }),
1175
+ gpuUsageWrap: {
1176
+ mb: 1.25,
1177
+ },
1178
+ gpuFooterRow: {
1179
+ display: 'flex',
1180
+ justifyContent: 'space-between',
1181
+ alignItems: 'center',
1182
+ },
1183
+ gpuFreeText: {
1184
+ fontWeight: 'bold',
1185
+ },
1186
+ centeredCaption: {
1187
+ display: 'block',
1188
+ textAlign: 'center',
1189
+ },
1190
+ centeredCaptionWithMargin: {
1191
+ display: 'block',
1192
+ textAlign: 'center',
1193
+ mt: 0.5,
1194
+ },
1195
+ dataProcessingGrid: {
1196
+ flex: 1,
1197
+ minHeight: 0,
1198
+ flexWrap: 'wrap',
1199
+ alignItems: 'stretch',
1200
+ },
1201
+ primaryPaneItem: {
1202
+ display: 'flex',
1203
+ flexDirection: 'column',
1204
+ minHeight: 0,
1205
+ overflow: 'visible',
1206
+ },
1207
+ primaryPaneContent: {
1208
+ flex: 1,
1209
+ overflow: 'visible',
1210
+ pr: { xs: 0, md: 1 },
1211
+ },
1212
+ secondaryPaneItem: {
1213
+ display: 'flex',
1214
+ flexDirection: 'column',
1215
+ overflow: 'visible',
1216
+ },
1217
+ secondaryPaneContent: {
1218
+ flex: 1,
1219
+ overflow: 'visible',
1220
+ pl: { xs: 0, md: 1 },
1221
+ },
1222
+ addRowButton: {
1223
+ mb: { xs: 2, sm: 3 },
1224
+ },
1225
+ sectionInfoAlert: {
1226
+ mb: 2,
1227
+ },
1228
+ recentFilesBlock: {
1229
+ mt: 1,
1230
+ },
1231
+ responsiveGrid: {
1232
+ height: 'auto',
1233
+ flexWrap: 'wrap',
1234
+ },
1235
+ formControlMarginBottom: {
1236
+ mb: 2,
1237
+ },
1238
+ fieldMarginBottom: {
1239
+ mb: 2,
1240
+ },
1241
+ fieldMarginBottomLarge: {
1242
+ mb: 3,
1243
+ },
1244
+ accordionMarginBottom: {
1245
+ mb: 2,
1246
+ },
1247
+ sliderRow: {
1248
+ display: 'flex',
1249
+ alignItems: 'center',
1250
+ gap: 2,
1251
+ },
1252
+ sliderFlexGrow: {
1253
+ flex: 1,
1254
+ },
1255
+ sliderInputSmall: {
1256
+ width: { xs: '72px', sm: '80px' },
1257
+ },
1258
+ sliderInputMedium: {
1259
+ width: { xs: '88px', sm: '100px' },
1260
+ },
1261
+ trainingActionRow: {
1262
+ display: 'flex',
1263
+ gap: { xs: 1.25, sm: 2 },
1264
+ flexDirection: { xs: 'column', sm: 'row' },
1265
+ },
1266
+ actionButtonFlexGrow: {
1267
+ flex: 1,
1268
+ },
1269
+ mediumWeightBodyText: {
1270
+ fontWeight: 500,
1271
+ },
1272
+ trainingMonitorWrap: {
1273
+ flex: 1,
1274
+ display: 'flex',
1275
+ flexDirection: 'column',
1276
+ },
1277
+ trainingMonitorHeaderRow: {
1278
+ display: 'flex',
1279
+ alignItems: 'center',
1280
+ justifyContent: 'space-between',
1281
+ mb: 1,
1282
+ },
1283
+ trainingMonitorTitle: {
1284
+ mb: 0,
1285
+ },
1286
+ trainingMonitorStatusInline: (muiTheme) => ({
1287
+ display: 'flex',
1288
+ alignItems: 'center',
1289
+ gap: 0.75,
1290
+ px: 1,
1291
+ py: 0.35,
1292
+ borderRadius: 999,
1293
+ border: '1px solid',
1294
+ borderColor: 'divider',
1295
+ backgroundColor: muiTheme.palette.mode === 'dark'
1296
+ ? 'rgba(10, 15, 23, 0.7)'
1297
+ : 'rgba(255, 255, 255, 0.9)',
1298
+ }),
1299
+ trainingMonitorStatusDot: (status, animate = false) => ({
1300
+ width: 8,
1301
+ height: 8,
1302
+ borderRadius: '50%',
1303
+ bgcolor: status === 'live'
1304
+ ? 'success.main'
1305
+ : status === 'error'
1306
+ ? 'error.main'
1307
+ : status === 'complete'
1308
+ ? 'primary.main'
1309
+ : 'text.secondary',
1310
+ animation: animate ? 'pulse 2s infinite' : 'none',
1311
+ '@keyframes pulse': {
1312
+ '0%': { opacity: 1 },
1313
+ '50%': { opacity: 0.5 },
1314
+ '100%': { opacity: 1 },
1315
+ },
1316
+ }),
1317
+ trainingMonitorStatusText: (status) => ({
1318
+ color: status === 'live'
1319
+ ? 'success.main'
1320
+ : status === 'error'
1321
+ ? 'error.main'
1322
+ : status === 'complete'
1323
+ ? 'primary.main'
1324
+ : 'text.secondary',
1325
+ fontWeight: 600,
1326
+ letterSpacing: '0.03em',
1327
+ textTransform: 'uppercase',
1328
+ fontSize: { xs: '0.64rem', sm: '0.7rem' },
1329
+ lineHeight: 1,
1330
+ }),
1331
+ generationModelRow: {
1332
+ display: 'flex',
1333
+ alignItems: 'center',
1334
+ gap: { xs: 1, sm: 2 },
1335
+ mb: 2,
1336
+ },
1337
+ refreshModelsButton: {
1338
+ minWidth: 40,
1339
+ },
1340
+ durationRow: {
1341
+ display: 'flex',
1342
+ alignItems: 'center',
1343
+ gap: { xs: 1, sm: 2 },
1344
+ mb: 2,
1345
+ },
1346
+ generatingWrap: {
1347
+ mb: 3,
1348
+ },
1349
+ generatingHeader: {
1350
+ display: 'flex',
1351
+ alignItems: 'center',
1352
+ mb: 1,
1353
+ },
1354
+ generatingSpinner: {
1355
+ mr: 1,
1356
+ },
1357
+ generatingProgress: {
1358
+ height: 8,
1359
+ borderRadius: 4,
1360
+ },
1361
+ generatingHint: {
1362
+ mt: 1,
1363
+ display: 'block',
1364
+ },
1365
+ generateButton: {
1366
+ mb: 2,
1367
+ },
1368
+ warningAlertTop: {
1369
+ mt: 2,
1370
+ },
1371
+ sectionCardHeader: {
1372
+ display: 'flex',
1373
+ alignItems: 'center',
1374
+ gap: 1,
1375
+ mb: 1.5,
1376
+ },
1377
+ sectionCardIcon: {
1378
+ display: 'inline-flex',
1379
+ color: 'text.primary',
1380
+ lineHeight: 0,
1381
+ },
1382
+ sectionCardTitle: {
1383
+ fontWeight: 500,
1384
+ },
1385
+ selectedModelCard: (muiTheme) => {
1386
+ const isDark = muiTheme.palette.mode === 'dark';
1387
+ return {
1388
+ p: { xs: 1.5, sm: 2 },
1389
+ mb: 2,
1390
+ boxShadow: isDark
1391
+ ? '0 14px 28px rgba(4, 8, 14, 0.44)'
1392
+ : '0 14px 26px rgba(15, 23, 42, 0.1)',
1393
+ borderRadius: 2.5,
1394
+ border: isDark
1395
+ ? '1px solid rgba(194, 207, 228, 0.16)'
1396
+ : '1px solid rgba(15, 23, 42, 0.12)',
1397
+ background: isDark
1398
+ ? 'linear-gradient(160deg, rgba(17, 24, 37, 0.96) 0%, rgba(13, 20, 31, 0.92) 100%)'
1399
+ : 'linear-gradient(160deg, rgba(255, 255, 255, 0.98) 0%, rgba(245, 250, 255, 0.98) 100%)',
1400
+ '&:hover': {
1401
+ boxShadow: isDark
1402
+ ? '0 20px 34px rgba(4, 8, 14, 0.56)'
1403
+ : '0 20px 34px rgba(15, 23, 42, 0.14)',
1404
+ transform: 'translateY(-1px)',
1405
+ transition: 'all 0.3s ease',
1406
+ },
1407
+ transition: 'all 0.3s ease',
1408
+ };
1409
+ },
1410
+ boldBodyText: {
1411
+ fontWeight: 'bold',
1412
+ },
1413
+ selectedModelMetaText: {
1414
+ display: 'block',
1415
+ mt: 0.5,
1416
+ },
1417
+ unwrappedInfoWrap: {
1418
+ mt: 2,
1419
+ },
1420
+ dialogBodyText: {
1421
+ mt: 3,
1422
+ },
1423
+ dialogErrorText: {
1424
+ mt: 2,
1425
+ },
1426
+ modeToggleButton: (muiTheme) => {
1427
+ const isDark = muiTheme.palette.mode === 'dark';
1428
+ return {
1429
+ position: 'fixed',
1430
+ left: { xs: 12, sm: 16 },
1431
+ bottom: { xs: 58, sm: 66 },
1432
+ width: { xs: 38, sm: 42 },
1433
+ height: { xs: 38, sm: 42 },
1434
+ zIndex: 1350,
1435
+ border: isDark
1436
+ ? '1px solid rgba(194, 207, 228, 0.22)'
1437
+ : '1px solid rgba(15, 23, 42, 0.16)',
1438
+ background: isDark
1439
+ ? 'linear-gradient(145deg, rgba(18, 25, 38, 0.96) 0%, rgba(12, 19, 30, 0.96) 100%)'
1440
+ : 'linear-gradient(145deg, rgba(255, 255, 255, 0.98) 0%, rgba(245, 250, 255, 0.98) 100%)',
1441
+ color: isDark ? 'primary.light' : 'primary.main',
1442
+ boxShadow: isDark
1443
+ ? '0 14px 24px rgba(4, 8, 14, 0.5)'
1444
+ : '0 14px 24px rgba(15, 23, 42, 0.14)',
1445
+ '&:hover': {
1446
+ background: isDark
1447
+ ? 'linear-gradient(145deg, rgba(20, 28, 42, 1) 0%, rgba(14, 22, 34, 1) 100%)'
1448
+ : 'linear-gradient(145deg, rgba(244, 250, 255, 1) 0%, rgba(236, 245, 252, 1) 100%)',
1449
+ transform: 'translateY(-1px)',
1450
+ boxShadow: isDark
1451
+ ? '0 18px 28px rgba(4, 8, 14, 0.6)'
1452
+ : '0 18px 28px rgba(15, 23, 42, 0.18)',
1453
+ },
1454
+ };
1455
+ },
1456
+ infoButton: (muiTheme) => {
1457
+ const isDark = muiTheme.palette.mode === 'dark';
1458
+ return {
1459
+ position: 'fixed',
1460
+ left: { xs: 12, sm: 16 },
1461
+ bottom: { xs: 12, sm: 16 },
1462
+ width: { xs: 38, sm: 42 },
1463
+ height: { xs: 38, sm: 42 },
1464
+ zIndex: 1350,
1465
+ border: isDark
1466
+ ? '1px solid rgba(194, 207, 228, 0.22)'
1467
+ : '1px solid rgba(15, 23, 42, 0.16)',
1468
+ background: isDark
1469
+ ? 'linear-gradient(145deg, rgba(18, 25, 38, 0.96) 0%, rgba(12, 19, 30, 0.96) 100%)'
1470
+ : 'linear-gradient(145deg, rgba(255, 255, 255, 0.98) 0%, rgba(245, 250, 255, 0.98) 100%)',
1471
+ color: isDark ? 'primary.light' : 'primary.main',
1472
+ boxShadow: isDark
1473
+ ? '0 14px 24px rgba(4, 8, 14, 0.5)'
1474
+ : '0 14px 24px rgba(15, 23, 42, 0.14)',
1475
+ '&:hover': {
1476
+ background: isDark
1477
+ ? 'linear-gradient(145deg, rgba(20, 28, 42, 1) 0%, rgba(14, 22, 34, 1) 100%)'
1478
+ : 'linear-gradient(145deg, rgba(244, 250, 255, 1) 0%, rgba(236, 245, 252, 1) 100%)',
1479
+ transform: 'translateY(-1px)',
1480
+ boxShadow: isDark
1481
+ ? '0 18px 28px rgba(4, 8, 14, 0.6)'
1482
+ : '0 18px 28px rgba(15, 23, 42, 0.18)',
1483
+ },
1484
+ };
1485
+ },
1486
+ infoDialogTitleRow: {
1487
+ display: 'inline-flex',
1488
+ alignItems: 'center',
1489
+ gap: 1,
1490
+ },
1491
+ infoDialogIntro: {
1492
+ mt: 1,
1493
+ color: 'text.secondary',
1494
+ },
1495
+ infoDialogSectionTitle: {
1496
+ mt: 2.25,
1497
+ mb: 1,
1498
+ color: 'text.primary',
1499
+ fontWeight: 700,
1500
+ letterSpacing: '0.03em',
1501
+ textTransform: 'uppercase',
1502
+ fontSize: { xs: '0.7rem', sm: '0.74rem' },
1503
+ },
1504
+ infoDialogActionStack: {
1505
+ mt: 0.25,
1506
+ display: 'flex',
1507
+ flexDirection: 'column',
1508
+ gap: 1,
1509
+ },
1510
+ infoDocButton: {
1511
+ justifyContent: 'flex-start',
1512
+ },
1513
+ };
1514
+
1515
+ export const tabPanelStyles = {
1516
+ root: {
1517
+ p: { xs: 1.25, md: 2 },
1518
+ background: 'transparent',
1519
+ flex: 1,
1520
+ display: 'flex',
1521
+ flexDirection: 'column',
1522
+ minHeight: 0,
1523
+ overflow: 'visible',
1524
+ },
1525
+ };
1526
+
1527
+ export const audioUploadRowStyles = {
1528
+ card: (muiTheme) => {
1529
+ const isDark = muiTheme.palette.mode === 'dark';
1530
+ return {
1531
+ mb: { xs: 1.5, sm: 2 },
1532
+ boxShadow: isDark
1533
+ ? '0 12px 24px rgba(4, 8, 14, 0.34)'
1534
+ : '0 12px 24px rgba(15, 23, 42, 0.1)',
1535
+ borderRadius: 2.2,
1536
+ border: isDark
1537
+ ? '1px solid rgba(194, 207, 228, 0.15)'
1538
+ : '1px solid rgba(15, 23, 42, 0.12)',
1539
+ background: isDark
1540
+ ? 'linear-gradient(160deg, rgba(17, 24, 37, 0.96) 0%, rgba(13, 20, 31, 0.92) 100%)'
1541
+ : 'linear-gradient(160deg, rgba(255, 255, 255, 0.99) 0%, rgba(245, 250, 255, 0.98) 100%)',
1542
+ '&:hover': {
1543
+ boxShadow: isDark
1544
+ ? '0 16px 30px rgba(4, 8, 14, 0.46)'
1545
+ : '0 16px 30px rgba(15, 23, 42, 0.14)',
1546
+ transform: 'translateY(-1px)',
1547
+ transition: 'all 0.3s ease',
1548
+ },
1549
+ transition: 'all 0.3s ease',
1550
+ };
1551
+ },
1552
+ cardContent: {
1553
+ p: { xs: 1.5, sm: 2 },
1554
+ '&:last-child': {
1555
+ pb: { xs: 1.5, sm: 2 },
1556
+ },
1557
+ },
1558
+ gridSpacing: { xs: 1.5, sm: 2 },
1559
+ uploadDropZone: (isDragActive) => (muiTheme) => {
1560
+ const isDark = muiTheme.palette.mode === 'dark';
1561
+ return {
1562
+ border: isDark
1563
+ ? '1.5px dashed rgba(194, 207, 228, 0.35)'
1564
+ : '1.5px dashed rgba(100, 116, 139, 0.34)',
1565
+ borderRadius: 2,
1566
+ p: { xs: 1.5, sm: 2 },
1567
+ textAlign: 'center',
1568
+ cursor: 'pointer',
1569
+ '&:hover': {
1570
+ borderColor: 'primary.main',
1571
+ backgroundColor: isDark
1572
+ ? 'rgba(53, 194, 212, 0.08)'
1573
+ : 'rgba(20, 151, 168, 0.08)',
1574
+ },
1575
+ backgroundColor: isDragActive
1576
+ ? (isDark ? 'rgba(53, 194, 212, 0.12)' : 'rgba(20, 151, 168, 0.12)')
1577
+ : (isDark ? 'rgba(10, 15, 23, 0.8)' : 'rgba(255, 255, 255, 0.9)'),
1578
+ };
1579
+ },
1580
+ hiddenInput: {
1581
+ display: 'none',
1582
+ },
1583
+ audioPreview: {
1584
+ width: '100%',
1585
+ marginTop: 8,
1586
+ },
1587
+ deleteGridItem: {
1588
+ display: 'flex',
1589
+ justifyContent: { xs: 'flex-end', sm: 'center' },
1590
+ },
1591
+ deleteIconButton: {
1592
+ alignSelf: { xs: 'center', sm: 'flex-start' },
1593
+ },
1594
+ };
1595
+
1596
+ export const generatedFragmentsWindowStyles = {
1597
+ rootPaper: (muiTheme) => {
1598
+ const isDark = muiTheme.palette.mode === 'dark';
1599
+ return {
1600
+ p: 2,
1601
+ height: 240,
1602
+ display: 'flex',
1603
+ flexDirection: 'column',
1604
+ borderRadius: 2.5,
1605
+ borderColor: isDark
1606
+ ? 'rgba(194, 207, 228, 0.16)'
1607
+ : 'rgba(15, 23, 42, 0.12)',
1608
+ background: isDark
1609
+ ? 'linear-gradient(160deg, rgba(17, 24, 37, 0.94) 0%, rgba(13, 20, 31, 0.9) 100%)'
1610
+ : 'linear-gradient(160deg, rgba(255, 255, 255, 0.99) 0%, rgba(245, 250, 255, 0.98) 100%)',
1611
+ boxShadow: isDark
1612
+ ? '0 14px 28px rgba(4, 8, 14, 0.42)'
1613
+ : '0 14px 28px rgba(15, 23, 42, 0.1)',
1614
+ };
1615
+ },
1616
+ headerRow: {
1617
+ display: 'flex',
1618
+ justifyContent: 'space-between',
1619
+ alignItems: 'center',
1620
+ mb: 2,
1621
+ },
1622
+ titleRow: {
1623
+ display: 'flex',
1624
+ alignItems: 'center',
1625
+ gap: 1,
1626
+ minWidth: 0,
1627
+ },
1628
+ titleIcon: {
1629
+ display: 'inline-flex',
1630
+ color: 'text.primary',
1631
+ lineHeight: 0,
1632
+ },
1633
+ titleText: {
1634
+ fontWeight: 500,
1635
+ },
1636
+ countText: {
1637
+ fontWeight: 600,
1638
+ minWidth: 20,
1639
+ textAlign: 'right',
1640
+ },
1641
+ emptyState: {
1642
+ display: 'flex',
1643
+ alignItems: 'center',
1644
+ justifyContent: 'center',
1645
+ height: '100%',
1646
+ color: 'text.secondary',
1647
+ },
1648
+ listRoot: (muiTheme) => {
1649
+ const isDark = muiTheme.palette.mode === 'dark';
1650
+ return {
1651
+ flex: 1,
1652
+ overflow: 'auto',
1653
+ maxHeight: 180,
1654
+ '& .MuiListItem-root': {
1655
+ border: '1px solid',
1656
+ borderColor: isDark
1657
+ ? 'rgba(194, 207, 228, 0.16)'
1658
+ : 'rgba(15, 23, 42, 0.12)',
1659
+ borderRadius: 1.5,
1660
+ mb: 1,
1661
+ backgroundColor: isDark
1662
+ ? 'rgba(12, 18, 28, 0.62)'
1663
+ : 'rgba(248, 251, 255, 0.9)',
1664
+ '&:last-child': {
1665
+ mb: 0,
1666
+ },
1667
+ },
1668
+ };
1669
+ },
1670
+ listItem: {
1671
+ display: 'flex',
1672
+ flexDirection: 'column',
1673
+ alignItems: 'stretch',
1674
+ py: 1,
1675
+ },
1676
+ fragmentRow: {
1677
+ display: 'flex',
1678
+ justifyContent: 'space-between',
1679
+ alignItems: 'flex-start',
1680
+ mb: 1,
1681
+ },
1682
+ fragmentMeta: {
1683
+ flex: 1,
1684
+ minWidth: 0,
1685
+ },
1686
+ fragmentPrompt: {
1687
+ fontWeight: 'bold',
1688
+ overflow: 'hidden',
1689
+ textOverflow: 'ellipsis',
1690
+ display: '-webkit-box',
1691
+ WebkitLineClamp: 2,
1692
+ WebkitBoxOrient: 'vertical',
1693
+ },
1694
+ fragmentActions: {
1695
+ display: 'flex',
1696
+ gap: 1,
1697
+ flexShrink: 0,
1698
+ },
1699
+ playPauseButton: (isPlaying) => (muiTheme) => ({
1700
+ border: '1px solid',
1701
+ borderColor: isPlaying
1702
+ ? 'primary.main'
1703
+ : (muiTheme.palette.mode === 'dark' ? 'rgba(194, 207, 228, 0.22)' : 'rgba(15, 23, 42, 0.2)'),
1704
+ }),
1705
+ hiddenAudio: {
1706
+ display: 'none',
1707
+ },
1708
+ };
1709
+
1710
+ export const trainingMonitorStyles = {
1711
+ rootPaper: (muiTheme) => {
1712
+ const isDark = muiTheme.palette.mode === 'dark';
1713
+ return {
1714
+ p: 3,
1715
+ mb: 2,
1716
+ flex: 1,
1717
+ display: 'flex',
1718
+ flexDirection: 'column',
1719
+ boxShadow: isDark
1720
+ ? '0 16px 30px rgba(4, 8, 14, 0.48)'
1721
+ : '0 16px 30px rgba(15, 23, 42, 0.1)',
1722
+ borderRadius: 2.5,
1723
+ border: isDark
1724
+ ? '1px solid rgba(194, 207, 228, 0.16)'
1725
+ : '1px solid rgba(15, 23, 42, 0.12)',
1726
+ background: isDark
1727
+ ? 'linear-gradient(160deg, rgba(17, 24, 37, 0.96) 0%, rgba(13, 20, 31, 0.92) 100%)'
1728
+ : 'linear-gradient(160deg, rgba(255, 255, 255, 0.99) 0%, rgba(245, 250, 255, 0.98) 100%)',
1729
+ '&:hover': {
1730
+ boxShadow: isDark
1731
+ ? '0 22px 38px rgba(4, 8, 14, 0.58)'
1732
+ : '0 22px 38px rgba(15, 23, 42, 0.14)',
1733
+ transform: 'translateY(-1px)',
1734
+ transition: 'all 0.3s ease',
1735
+ },
1736
+ transition: 'all 0.3s ease',
1737
+ };
1738
+ },
1739
+ headerRow: {
1740
+ display: 'flex',
1741
+ justifyContent: 'space-between',
1742
+ alignItems: 'center',
1743
+ flexWrap: 'wrap',
1744
+ gap: 1,
1745
+ mb: 2,
1746
+ },
1747
+ headerTitleWrap: {
1748
+ display: 'flex',
1749
+ alignItems: 'center',
1750
+ gap: 1,
1751
+ },
1752
+ headerIcon: {
1753
+ display: 'inline-flex',
1754
+ color: 'text.primary',
1755
+ lineHeight: 0,
1756
+ },
1757
+ headerTitle: {
1758
+ fontWeight: 500,
1759
+ },
1760
+ statusInline: (muiTheme) => ({
1761
+ display: 'flex',
1762
+ alignItems: 'center',
1763
+ gap: 0.75,
1764
+ px: 1,
1765
+ py: 0.35,
1766
+ borderRadius: 999,
1767
+ border: '1px solid',
1768
+ borderColor: 'divider',
1769
+ backgroundColor: muiTheme.palette.mode === 'dark'
1770
+ ? 'rgba(10, 15, 23, 0.7)'
1771
+ : 'rgba(255, 255, 255, 0.9)',
1772
+ }),
1773
+ statusDot: (status, animate = false) => ({
1774
+ width: 8,
1775
+ height: 8,
1776
+ borderRadius: '50%',
1777
+ bgcolor: status === 'live'
1778
+ ? 'success.main'
1779
+ : status === 'error'
1780
+ ? 'error.main'
1781
+ : status === 'complete'
1782
+ ? 'primary.main'
1783
+ : 'text.secondary',
1784
+ animation: animate ? 'pulse 2s infinite' : 'none',
1785
+ '@keyframes pulse': {
1786
+ '0%': { opacity: 1 },
1787
+ '50%': { opacity: 0.5 },
1788
+ '100%': { opacity: 1 },
1789
+ },
1790
+ }),
1791
+ statusText: (status) => ({
1792
+ color: status === 'live'
1793
+ ? 'success.main'
1794
+ : status === 'error'
1795
+ ? 'error.main'
1796
+ : status === 'complete'
1797
+ ? 'primary.main'
1798
+ : 'text.secondary',
1799
+ fontWeight: 600,
1800
+ letterSpacing: '0.03em',
1801
+ textTransform: 'uppercase',
1802
+ fontSize: { xs: '0.64rem', sm: '0.7rem' },
1803
+ lineHeight: 1,
1804
+ }),
1805
+ progressSection: {
1806
+ mb: 2,
1807
+ },
1808
+ progressHeader: {
1809
+ display: 'flex',
1810
+ justifyContent: 'space-between',
1811
+ mb: 1,
1812
+ },
1813
+ progressBar: {
1814
+ height: 8,
1815
+ borderRadius: 4,
1816
+ },
1817
+ deviceSection: {
1818
+ mb: 2,
1819
+ },
1820
+ deviceInfo: {
1821
+ fontSize: { xs: '0.74rem', sm: '0.8rem' },
1822
+ mt: 0.5,
1823
+ },
1824
+ metricsGrid: {
1825
+ mb: 2,
1826
+ },
1827
+ lossSection: {
1828
+ mb: 2,
1829
+ },
1830
+ lossChartBox: {
1831
+ height: 200,
1832
+ width: '100%',
1833
+ },
1834
+ errorAlert: {
1835
+ mb: 2,
1836
+ },
1837
+ };
1838
+
1839
+ export const welcomePageStyles = {
1840
+ backdrop: (muiTheme) => {
1841
+ const isDark = muiTheme.palette.mode === 'dark';
1842
+ return {
1843
+ zIndex: 9999,
1844
+ background: isDark
1845
+ ? 'radial-gradient(1200px 600px at 8% -15%, rgba(53, 194, 212, 0.2), transparent 55%), radial-gradient(800px 500px at 92% 110%, rgba(83, 193, 138, 0.16), transparent 65%), linear-gradient(160deg, #090C12 0%, #0C1119 45%, #090D13 100%)'
1846
+ : 'radial-gradient(1200px 600px at 8% -15%, rgba(20, 151, 168, 0.16), transparent 55%), radial-gradient(800px 500px at 92% 110%, rgba(72, 171, 118, 0.14), transparent 65%), linear-gradient(160deg, #F4FAFF 0%, #EAF4FF 45%, #F8FCFF 100%)',
1847
+ display: 'flex',
1848
+ alignItems: 'center',
1849
+ justifyContent: 'center',
1850
+ p: { xs: 2, md: 4 },
1851
+ cursor: 'pointer',
1852
+ };
1853
+ },
1854
+ panel: (muiTheme) => {
1855
+ const isDark = muiTheme.palette.mode === 'dark';
1856
+ return {
1857
+ textAlign: 'center',
1858
+ width: 'min(920px, 100%)',
1859
+ border: isDark
1860
+ ? '1px solid rgba(194, 207, 228, 0.2)'
1861
+ : '1px solid rgba(15, 23, 42, 0.12)',
1862
+ borderRadius: 4,
1863
+ background: isDark
1864
+ ? 'linear-gradient(170deg, rgba(19, 27, 41, 0.95) 0%, rgba(12, 19, 31, 0.96) 100%)'
1865
+ : 'linear-gradient(170deg, rgba(255, 255, 255, 0.98) 0%, rgba(242, 248, 255, 0.98) 100%)',
1866
+ boxShadow: isDark
1867
+ ? '0 32px 56px rgba(4, 8, 14, 0.64)'
1868
+ : '0 24px 46px rgba(15, 23, 42, 0.2)',
1869
+ backdropFilter: 'blur(10px)',
1870
+ px: { xs: 3, md: 7 },
1871
+ py: { xs: 5, md: 6 },
1872
+ };
1873
+ },
1874
+ logo: (muiTheme) => {
1875
+ const isDark = muiTheme.palette.mode === 'dark';
1876
+ return {
1877
+ width: { xs: 96, sm: 122 },
1878
+ height: { xs: 96, sm: 122 },
1879
+ backgroundImage: 'url(/fragmenta_icon_1024.png)',
1880
+ backgroundSize: 'cover',
1881
+ backgroundPosition: 'center',
1882
+ borderRadius: 3,
1883
+ border: isDark
1884
+ ? '1px solid rgba(194, 207, 228, 0.24)'
1885
+ : '1px solid rgba(15, 23, 42, 0.12)',
1886
+ boxShadow: isDark
1887
+ ? '0 16px 28px rgba(4, 8, 14, 0.45)'
1888
+ : '0 12px 24px rgba(15, 23, 42, 0.22)',
1889
+ filter: isDark
1890
+ ? 'drop-shadow(0 8px 16px rgba(0, 0, 0, 0.4))'
1891
+ : 'drop-shadow(0 6px 12px rgba(15, 23, 42, 0.22))',
1892
+ mx: 'auto',
1893
+ mb: 1.5,
1894
+ };
1895
+ },
1896
+ title: {
1897
+ fontFamily: '"Bitcount Single", "IBM Plex Mono", "JetBrains Mono", "Space Mono", "Courier New", monospace',
1898
+ fontWeight: 400,
1899
+ color: 'text.primary',
1900
+ mb: 1,
1901
+ fontSize: { xs: '2.5rem', sm: '3.5rem', md: '4rem' },
1902
+ letterSpacing: '0.02em',
1903
+ },
1904
+ overline: {
1905
+ color: 'primary.main',
1906
+ letterSpacing: { xs: '0.12em', md: '0.18em' },
1907
+ fontWeight: 700,
1908
+ fontSize: { xs: '0.62rem', sm: '0.7rem' },
1909
+ },
1910
+ footer: {
1911
+ color: 'text.secondary',
1912
+ opacity: 0.6,
1913
+ fontSize: { xs: '0.64rem', sm: '0.7rem' },
1914
+ marginTop: 5,
1915
+ },
1916
+ version: {
1917
+ color: 'text.secondary',
1918
+ opacity: 0.6,
1919
+ fontSize: { xs: '0.64rem', sm: '0.7rem' },
1920
+ fontStyle: 'italic',
1921
+ },
1922
+ ctaButton: {
1923
+ mt: 3,
1924
+ mb: 2,
1925
+ px: { xs: 3.25, md: 4.5 },
1926
+ py: { xs: 1.2, md: 1.5 },
1927
+ borderRadius: 2,
1928
+ textTransform: 'none',
1929
+ fontSize: { xs: '0.98rem', sm: '1.05rem', md: '1.1rem' },
1930
+ fontWeight: 500,
1931
+ },
1932
+ };
1933
+
1934
+ export const checkpointManagerStyles = {
1935
+ root: {
1936
+ mt: 2,
1937
+ pt: 2,
1938
+ borderTop: '1px solid',
1939
+ borderColor: 'divider',
1940
+ },
1941
+ panelPaper: (muiTheme) => {
1942
+ const isDark = muiTheme.palette.mode === 'dark';
1943
+ return {
1944
+ p: 2,
1945
+ mb: 2,
1946
+ boxShadow: isDark
1947
+ ? '0 16px 30px rgba(4, 8, 14, 0.48)'
1948
+ : '0 16px 30px rgba(15, 23, 42, 0.1)',
1949
+ borderRadius: 2.5,
1950
+ border: isDark
1951
+ ? '1px solid rgba(194, 207, 228, 0.16)'
1952
+ : '1px solid rgba(15, 23, 42, 0.12)',
1953
+ background: isDark
1954
+ ? 'linear-gradient(160deg, rgba(17, 24, 37, 0.96) 0%, rgba(13, 20, 31, 0.92) 100%)'
1955
+ : 'linear-gradient(160deg, rgba(255, 255, 255, 0.99) 0%, rgba(245, 250, 255, 0.98) 100%)',
1956
+ '&:hover': {
1957
+ boxShadow: isDark
1958
+ ? '0 22px 38px rgba(4, 8, 14, 0.58)'
1959
+ : '0 22px 38px rgba(15, 23, 42, 0.14)',
1960
+ transform: 'translateY(-1px)',
1961
+ transition: 'all 0.3s ease',
1962
+ },
1963
+ transition: 'all 0.3s ease',
1964
+ };
1965
+ },
1966
+ checkpointsList: {
1967
+ mt: 1,
1968
+ display: 'grid',
1969
+ gap: 1,
1970
+ },
1971
+ checkpointCard: (muiTheme) => {
1972
+ const isDark = muiTheme.palette.mode === 'dark';
1973
+ return {
1974
+ mb: 1,
1975
+ p: 1.25,
1976
+ boxShadow: isDark
1977
+ ? '0 10px 20px rgba(4, 8, 14, 0.34)'
1978
+ : '0 10px 18px rgba(15, 23, 42, 0.08)',
1979
+ borderRadius: 1.75,
1980
+ border: isDark
1981
+ ? '1px solid rgba(194, 207, 228, 0.16)'
1982
+ : '1px solid rgba(15, 23, 42, 0.12)',
1983
+ background: isDark
1984
+ ? 'linear-gradient(180deg, rgba(18, 25, 38, 0.98) 0%, rgba(15, 22, 34, 0.96) 100%)'
1985
+ : 'linear-gradient(180deg, rgba(255, 255, 255, 0.99) 0%, rgba(248, 251, 255, 0.99) 100%)',
1986
+ '&:hover': {
1987
+ boxShadow: isDark
1988
+ ? '0 14px 26px rgba(4, 8, 14, 0.44)'
1989
+ : '0 16px 28px rgba(15, 23, 42, 0.12)',
1990
+ transform: 'translateY(-1px)',
1991
+ transition: 'all 0.2s ease',
1992
+ },
1993
+ transition: 'all 0.2s ease',
1994
+ };
1995
+ },
1996
+ checkpointRow: {
1997
+ display: 'flex',
1998
+ alignItems: 'center',
1999
+ justifyContent: 'space-between',
2000
+ gap: 1,
2001
+ },
2002
+ checkpointInfo: {
2003
+ flex: 1,
2004
+ },
2005
+ checkpointName: {
2006
+ fontWeight: 600,
2007
+ display: 'flex',
2008
+ alignItems: 'center',
2009
+ flexWrap: 'wrap',
2010
+ gap: 0.75,
2011
+ },
2012
+ unwrappedChip: {
2013
+ fontSize: '0.7rem',
2014
+ },
2015
+ emptyText: {
2016
+ display: 'block',
2017
+ mt: 0.5,
2018
+ },
2019
+ metaNext: {
2020
+ ml: 1,
2021
+ },
2022
+ actions: {
2023
+ display: 'flex',
2024
+ gap: 1,
2025
+ flexWrap: 'wrap',
2026
+ justifyContent: 'flex-end',
2027
+ },
2028
+ errorAlert: {
2029
+ mt: 2,
2030
+ },
2031
+ deleteDialogText: {
2032
+ mt: 1.5,
2033
+ },
2034
+ snackbarAlert: {
2035
+ width: '100%',
2036
+ },
2037
+ };
2038
+
2039
+ export const hfAuthDialogStyles = {
2040
+ checkingBox: {
2041
+ display: 'flex',
2042
+ flexDirection: 'column',
2043
+ alignItems: 'center',
2044
+ my: 4,
2045
+ },
2046
+ checkingProgress: {
2047
+ mb: 2,
2048
+ },
2049
+ authStepBox: {
2050
+ mt: 2,
2051
+ },
2052
+ downloadStepBox: {
2053
+ mt: 2,
2054
+ textAlign: 'center',
2055
+ },
2056
+ downloadProgress: {
2057
+ mb: 2,
2058
+ height: 10,
2059
+ borderRadius: 5,
2060
+ },
2061
+ successStepBox: {
2062
+ mt: 2,
2063
+ textAlign: 'center',
2064
+ },
2065
+ stepper: {
2066
+ mb: 4,
2067
+ },
2068
+ errorAlert: {
2069
+ mb: 2,
2070
+ },
2071
+ loginSpinnerSize: 24,
2072
+ };
2073
+
2074
+ export const modelUnwrapButtonStyles = {
2075
+ root: {
2076
+ mt: 1,
2077
+ },
2078
+ result: {
2079
+ mt: 0.5,
2080
+ },
2081
+ error: {
2082
+ color: '#DB5044',
2083
+ mt: 0.5,
2084
+ },
2085
+ };
2086
+
2087
+ export const lossChartStyles = {
2088
+ padding: { top: 10, right: 16, bottom: 28, left: 44 },
2089
+ colors: {
2090
+ grid: '#2B3446',
2091
+ axis: '#9DA9BC',
2092
+ line: '#35C2D4',
2093
+ point: '#73D7E3',
2094
+ tooltipBg: '#111826',
2095
+ tooltipBorder: 'rgba(194, 207, 228, 0.2)',
2096
+ tooltipText: '#E8EDF5',
2097
+ },
2098
+ svg: {
2099
+ width: '100%',
2100
+ height: '100%',
2101
+ },
2102
+ axisFontSize: 10,
2103
+ tooltip: {
2104
+ width: 100,
2105
+ height: 34,
2106
+ rx: 4,
2107
+ textX: 6,
2108
+ timeY: 14,
2109
+ lossY: 28,
2110
+ },
2111
+ };
2112
+
2113
+ export default theme;
app/frontend/src/utils/format.js ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ export function formatDuration(seconds) {
2
+ const sec = Math.floor(seconds % 60);
3
+ const min = Math.floor((seconds / 60) % 60);
4
+ const hr = Math.floor(seconds / 3600);
5
+ return [hr, min, sec]
6
+ .map((v, i) => (i === 0 ? v : v.toString().padStart(2, '0')))
7
+ .join(':');
8
+ }
app/frontend/vite.config.js ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { defineConfig } from 'vite';
2
+ import react from '@vitejs/plugin-react';
3
+
4
+ export default defineConfig({
5
+ plugins: [react({ include: /\.(js|jsx)$/ })],
6
+ esbuild: {
7
+ loader: 'jsx',
8
+ include: /src\/.*\.jsx?$/,
9
+ exclude: [],
10
+ },
11
+ optimizeDeps: {
12
+ esbuildOptions: {
13
+ loader: { '.js': 'jsx' },
14
+ },
15
+ },
16
+ build: {
17
+ outDir: 'build',
18
+ sourcemap: false,
19
+ },
20
+ server: {
21
+ port: 3000,
22
+ proxy: {
23
+ '/api': 'http://localhost:5001',
24
+ },
25
+ },
26
+ });
docker-entrypoint.sh CHANGED
@@ -24,13 +24,22 @@ echo "[startup] Python: $(python --version 2>&1)"
24
  echo "[startup] Working dir: $(pwd)"
25
  echo "[startup] User: $(whoami) (uid=$(id -u))"
26
 
27
- # Quick smoke-test: can we import the core deps? (avoid CUDA init — it's slow)
28
- echo "[startup] Checking core Python imports…"
29
  python -c "
30
  import flask, flask_cors, torch, torchaudio, soundfile
31
  print(f' Flask {flask.__version__} | torch {torch.__version__}')
32
  print(f' CUDA build: {torch.version.cuda or \"none\"}')
33
- print(' GPU detection deferred to first request (faster startup)')
 
 
 
 
 
 
 
 
 
34
  " || {
35
  echo "ERROR: Core Python imports failed. Check requirements."
36
  exit 1
 
24
  echo "[startup] Working dir: $(pwd)"
25
  echo "[startup] User: $(whoami) (uid=$(id -u))"
26
 
27
+ # Smoke-test: verify core imports and detect GPU
28
+ echo "[startup] Checking core Python imports and GPU…"
29
  python -c "
30
  import flask, flask_cors, torch, torchaudio, soundfile
31
  print(f' Flask {flask.__version__} | torch {torch.__version__}')
32
  print(f' CUDA build: {torch.version.cuda or \"none\"}')
33
+ if torch.cuda.is_available():
34
+ name = torch.cuda.get_device_name(0)
35
+ vram = torch.cuda.get_device_properties(0).total_memory / 1024**3
36
+ print(f' GPU: {name} ({vram:.1f} GB VRAM) — CUDA ready')
37
+ else:
38
+ print(' No CUDA GPU detected — will run on CPU')
39
+ print(' If you have an NVIDIA GPU, check:')
40
+ print(' 1. Run \"docker compose up\" (not \"docker-compose up\")')
41
+ print(' 2. Docker Desktop is using the WSL2 backend')
42
+ print(' 3. \"nvidia-smi\" works on the host')
43
  " || {
44
  echo "ERROR: Core Python imports failed. Check requirements."
45
  exit 1
models/config/dataset-config.json CHANGED
@@ -2,10 +2,11 @@
2
  "dataset_type": "audio_dir",
3
  "datasets": [
4
  {
5
- "id": "fine_tune_data",
6
- "path": "/home/misagh/Documents/Fragmenta_Desktop/data",
7
  "custom_metadata_module": "custom_metadata"
8
  }
9
  ],
10
- "random_crop": true
 
11
  }
 
2
  "dataset_type": "audio_dir",
3
  "datasets": [
4
  {
5
+ "id": "custom_dataset",
6
+ "path": "/home/misagh/Documents/Fragmenta_Revised/data",
7
  "custom_metadata_module": "custom_metadata"
8
  }
9
  ],
10
+ "random_crop": true,
11
+ "drop_last": true
12
  }
requirements.txt CHANGED
@@ -1,8 +1,5 @@
1
  --extra-index-url https://download.pytorch.org/whl/cu128
2
 
3
- PyQt6
4
- PyQt6-WebEngine>=6.4.0
5
-
6
  Flask>=2.3.0
7
  Flask-CORS>=4.0.0
8
  requests>=2.28.0
@@ -15,32 +12,23 @@ diffusers>=0.20.0
15
  accelerate>=0.20.0
16
  peft>=0.4.0
17
  datasets>=2.14.0
18
- huggingface-hub>=0.16.0
19
- flash-attn>=2.8.3
20
 
21
  librosa>=0.10.0
22
  soundfile>=0.12.0
23
  scipy>=1.10.0
24
- numpy==1.23.5
25
-
26
- gradio>=3.40.0
27
- matplotlib>=3.7.0
28
- plotly>=5.15.0
29
 
30
- setuptools<82
31
  tqdm>=4.65.0
32
- psutil>=5.9.0
33
- wandb>=0.15.0
34
  omegaconf>=2.3.0
35
- click>=8.1.0
 
 
36
 
37
- pytest>=7.4.0
38
- black>=23.0.0
39
- isort>=5.12.0
40
- flake8>=6.0.0
41
- Pillow>=9.0.0
42
 
43
- fastapi>=0.100.0
44
- uvicorn>=0.23.0
45
- python-dotenv>=1.0.0
46
- setproctitle>=1.3.0
 
1
  --extra-index-url https://download.pytorch.org/whl/cu128
2
 
 
 
 
3
  Flask>=2.3.0
4
  Flask-CORS>=4.0.0
5
  requests>=2.28.0
 
12
  accelerate>=0.20.0
13
  peft>=0.4.0
14
  datasets>=2.14.0
15
+ huggingface-hub>=0.16.0
 
16
 
17
  librosa>=0.10.0
18
  soundfile>=0.12.0
19
  scipy>=1.10.0
20
+ numpy==1.23.5
 
 
 
 
21
 
22
+ setuptools<70
23
  tqdm>=4.65.0
24
+ psutil>=5.9.0
 
25
  omegaconf>=2.3.0
26
+ click>=8.1.0
27
+ Pillow>=9.0.0
28
+ python-dotenv>=1.0.0
29
 
30
+ pywebview>=4.4.1
31
+ Pycairo ; sys_platform == 'linux'
32
+ PyGObject<3.49 ; sys_platform == 'linux'
 
 
33
 
34
+ flash-attn>=2.8.3
 
 
 
stable-audio-tools/custom_metadata.py CHANGED
@@ -6,25 +6,25 @@ logging.basicConfig(level=logging.WARNING)
6
 
7
  def _load_metadata():
8
  metadata = {}
9
- project_root = os.path.join(os.path.dirname(__file__), '..')
 
 
 
 
10
 
11
- # Primary location: <project_root>/data/metadata.json (where the UI saves it)
12
- primary_path = os.path.join(project_root, 'data', 'metadata.json')
13
- # Legacy fallback: <project_root>/app/backend/data/metadata.json
14
- legacy_path = os.path.join(project_root, 'app', 'backend', 'data', 'metadata.json')
15
-
16
- json_path = primary_path if os.path.isfile(primary_path) else legacy_path
17
 
18
  try:
19
- with open(json_path, 'r') as jsonfile:
20
  metadata_list = json.load(jsonfile)
21
  for item in metadata_list:
22
  metadata[item['file_name']] = item['prompt']
23
- except FileNotFoundError:
24
- logging.warning(f"Metadata file not found at {primary_path} or {legacy_path}")
25
  except Exception as e:
26
- logging.warning(f"Error loading metadata: {e}")
27
-
28
  return metadata
29
 
30
  def get_custom_metadata(info, audio):
 
6
 
7
  def _load_metadata():
8
  metadata = {}
9
+ here = os.path.dirname(__file__)
10
+ candidates = [
11
+ os.path.join(here, '..', 'data', 'metadata.json'), # current layout
12
+ os.path.join(here, '..', 'app', 'backend', 'data', 'metadata.json'), # legacy
13
+ ]
14
 
15
+ chosen = next((p for p in candidates if os.path.exists(p)), None)
16
+ if chosen is None:
17
+ logging.warning(f"Metadata file not found in any of: {candidates}")
18
+ return metadata
 
 
19
 
20
  try:
21
+ with open(chosen, 'r') as jsonfile:
22
  metadata_list = json.load(jsonfile)
23
  for item in metadata_list:
24
  metadata[item['file_name']] = item['prompt']
 
 
25
  except Exception as e:
26
+ logging.warning(f"Error loading metadata from {chosen}: {e}")
27
+
28
  return metadata
29
 
30
  def get_custom_metadata(info, audio):
stable-audio-tools/package-lock.json ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ {
2
+ "name": "stable-audio-tools",
3
+ "lockfileVersion": 3,
4
+ "requires": true,
5
+ "packages": {}
6
+ }
stable-audio-tools/stable_audio_tools/models/conditioners.py CHANGED
@@ -261,7 +261,7 @@ class CLAPAudioConditioner(Conditioner):
261
  # Convert to mono
262
  mono_audios = audios.mean(dim=1)
263
 
264
- with torch.amp.autocast(device_type=str(device) if str(device) in ('cuda', 'cpu') else 'cpu', enabled=False):
265
  audio_embedding = self.model.get_audio_embedding_from_data(mono_audios.float(), use_tensor=True)
266
 
267
  audio_embedding = audio_embedding.unsqueeze(1).to(device)
@@ -350,8 +350,11 @@ class T5Conditioner(Conditioner):
350
 
351
  self.model.eval()
352
 
353
- _autocast_device = 'cuda' if torch.cuda.is_available() else 'cpu'
354
- with torch.amp.autocast(device_type=_autocast_device, dtype=torch.float16) and torch.set_grad_enabled(self.enable_grad):
 
 
 
355
  embeddings = self.model(
356
  input_ids=input_ids, attention_mask=attention_mask
357
  )["last_hidden_state"]
 
261
  # Convert to mono
262
  mono_audios = audios.mean(dim=1)
263
 
264
+ with torch.amp.autocast("cuda", enabled=False):
265
  audio_embedding = self.model.get_audio_embedding_from_data(mono_audios.float(), use_tensor=True)
266
 
267
  audio_embedding = audio_embedding.unsqueeze(1).to(device)
 
350
 
351
  self.model.eval()
352
 
353
+ with torch.amp.autocast(
354
+ "cuda",
355
+ dtype=torch.float16,
356
+ enabled=str(device).startswith("cuda"),
357
+ ) and torch.set_grad_enabled(self.enable_grad):
358
  embeddings = self.model(
359
  input_ids=input_ids, attention_mask=attention_mask
360
  )["last_hidden_state"]
stable-audio-tools/stable_audio_tools/training/arc.py CHANGED
@@ -285,7 +285,7 @@ class ARCTrainingWrapper(pl.LightningModule):
285
  self.diffusion.pretransform.to(self.device)
286
 
287
  if not self.pre_encoded:
288
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu') and torch.set_grad_enabled(self.diffusion.pretransform.enable_grad):
289
  self.diffusion.pretransform.train(self.diffusion.pretransform.enable_grad)
290
 
291
  diffusion_input = self.diffusion.pretransform.encode(diffusion_input)
 
285
  self.diffusion.pretransform.to(self.device)
286
 
287
  if not self.pre_encoded:
288
+ with torch.cuda.amp.autocast() and torch.set_grad_enabled(self.diffusion.pretransform.enable_grad):
289
  self.diffusion.pretransform.train(self.diffusion.pretransform.enable_grad)
290
 
291
  diffusion_input = self.diffusion.pretransform.encode(diffusion_input)
stable-audio-tools/stable_audio_tools/training/diffusion.py CHANGED
@@ -120,7 +120,7 @@ class DiffusionUncondTrainingWrapper(pl.LightningModule):
120
  noised_inputs = diffusion_input * alphas + noise * sigmas
121
  targets = noise * alphas - diffusion_input * sigmas
122
 
123
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
124
  v = self.diffusion(noised_inputs, t)
125
 
126
  loss_info.update({
@@ -185,7 +185,7 @@ class DiffusionUncondDemoCallback(pl.Callback):
185
  noise = torch.randn([self.num_demos, module.diffusion.io_channels, demo_samples]).to(module.device)
186
 
187
  try:
188
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
189
  fakes = sample(module.diffusion_ema, noise, self.demo_steps, 0)
190
 
191
  if module.diffusion.pretransform is not None:
@@ -347,7 +347,7 @@ class DiffusionCondTrainingWrapper(pl.LightningModule):
347
  self.diffusion.pretransform.to(self.device)
348
 
349
  if not self.pre_encoded:
350
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu') and torch.set_grad_enabled(self.diffusion.pretransform.enable_grad):
351
  self.diffusion.pretransform.train(self.diffusion.pretransform.enable_grad)
352
 
353
  diffusion_input = self.diffusion.pretransform.encode(diffusion_input)
@@ -467,7 +467,7 @@ class DiffusionCondTrainingWrapper(pl.LightningModule):
467
 
468
  diffusion_input = reals
469
 
470
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu') and torch.no_grad():
471
  conditioning = self.diffusion.conditioner(metadata, self.device)
472
 
473
  # TODO: decide what to do with padding masks during validation
@@ -483,7 +483,7 @@ class DiffusionCondTrainingWrapper(pl.LightningModule):
483
  self.diffusion.pretransform.to(self.device)
484
 
485
  if not self.pre_encoded:
486
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu') and torch.no_grad():
487
  self.diffusion.pretransform.train(self.diffusion.pretransform.enable_grad)
488
 
489
  diffusion_input = self.diffusion.pretransform.encode(diffusion_input)
@@ -522,7 +522,7 @@ class DiffusionCondTrainingWrapper(pl.LightningModule):
522
  # if use_padding_mask:
523
  # extra_args["mask"] = padding_masks
524
 
525
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu') and torch.no_grad():
526
  output = self.diffusion(noised_inputs, t, cond=conditioning, cfg_dropout_prob = 0, **extra_args)
527
 
528
  val_loss = F.mse_loss(output, targets)
@@ -621,7 +621,7 @@ class DiffusionCondDemoCallback(pl.Callback):
621
 
622
  try:
623
  print("Getting conditioning")
624
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
625
  conditioning = module.diffusion.conditioner(demo_cond, module.device)
626
 
627
  cond_inputs = module.diffusion.get_conditioning_inputs(conditioning)
@@ -665,7 +665,7 @@ class DiffusionCondDemoCallback(pl.Callback):
665
 
666
  print(f"Generating demo for cfg scale {cfg_scale}")
667
 
668
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
669
  model = module.diffusion_ema.model if module.diffusion_ema is not None else module.diffusion.model
670
 
671
  if module.diffusion_objective == "v":
@@ -917,7 +917,7 @@ class DiffusionCondInpaintTrainingWrapper(pl.LightningModule):
917
 
918
  p.tick("setup")
919
 
920
- #with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
921
  conditioning = self.diffusion.conditioner(metadata, self.device)
922
 
923
  p.tick("conditioning")
@@ -983,7 +983,7 @@ class DiffusionCondInpaintTrainingWrapper(pl.LightningModule):
983
 
984
  extra_args = {}
985
 
986
- #with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
987
  p.tick("amp")
988
  output = self.diffusion(noised_inputs, t, cond=conditioning, cfg_dropout_prob = self.cfg_dropout_prob, **extra_args)
989
  p.tick("diffusion")
@@ -1127,7 +1127,7 @@ class DiffusionCondInpaintDemoCallback(pl.Callback):
1127
  model = module.diffusion_ema.model if module.diffusion_ema is not None else module.diffusion.model
1128
  print(f"Generating demo for cfg scale {cfg_scale}")
1129
 
1130
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
1131
  if module.diffusion_objective == "v":
1132
  fakes = sample(model, noise, self.demo_steps, 0, **cond_inputs, cfg_scale=cfg_scale, dist_shift=module.diffusion.dist_shift, batch_cfg=True)
1133
  elif module.diffusion_objective == "rectified_flow":
@@ -1274,7 +1274,7 @@ class DiffusionAutoencoderTrainingWrapper(pl.LightningModule):
1274
  noised_reals = reals * alphas + noise * sigmas
1275
  targets = noise * alphas - reals * sigmas
1276
 
1277
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
1278
  v = self.diffae.diffusion(noised_reals, t, input_concat_cond=latents)
1279
 
1280
  loss_info.update({
@@ -1354,7 +1354,7 @@ class DiffusionAutoencoderDemoCallback(pl.Callback):
1354
 
1355
  demo_reals = demo_reals.to(module.device)
1356
 
1357
- with torch.no_grad() and torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
1358
  latents = module.diffae_ema.ema_model.encode(encoder_input).float()
1359
  fakes = module.diffae_ema.ema_model.decode(latents, steps=self.demo_steps)
1360
 
@@ -1387,7 +1387,7 @@ class DiffusionAutoencoderDemoCallback(pl.Callback):
1387
  audio_spectrogram_image(reals_fakes))
1388
 
1389
  if module.diffae_ema.ema_model.pretransform is not None:
1390
- with torch.no_grad() and torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
1391
  initial_latents = module.diffae_ema.ema_model.pretransform.encode(encoder_input)
1392
  first_stage_fakes = module.diffae_ema.ema_model.pretransform.decode(initial_latents)
1393
  first_stage_fakes = rearrange(first_stage_fakes, 'b d n -> d (b n)')
@@ -1546,7 +1546,7 @@ class DiffusionPriorTrainingWrapper(pl.LightningModule):
1546
  source = self.diffusion.pretransform.encode(source)
1547
 
1548
  if self.diffusion.conditioner is not None:
1549
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
1550
  conditioning = self.diffusion.conditioner(metadata, self.device)
1551
  else:
1552
  conditioning = {}
@@ -1566,7 +1566,7 @@ class DiffusionPriorTrainingWrapper(pl.LightningModule):
1566
  noised_reals = reals * alphas + noise * sigmas
1567
  targets = noise * alphas - reals * sigmas
1568
 
1569
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
1570
 
1571
  conditioning['source'] = [source]
1572
 
@@ -1675,13 +1675,13 @@ class DiffusionPriorDemoCallback(pl.Callback):
1675
  encoder_input = demo_reals
1676
 
1677
  if module.diffusion.conditioner is not None:
1678
- with torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
1679
  conditioning_tensors = module.diffusion.conditioner(metadata, module.device)
1680
 
1681
  else:
1682
  conditioning_tensors = {}
1683
 
1684
- with torch.no_grad() and torch.amp.autocast('cuda' if torch.cuda.is_available() else 'cpu'):
1685
  if module.prior_type == PriorType.MonoToStereo and encoder_input.shape[1] > 1:
1686
  source = encoder_input.mean(dim=1, keepdim=True).repeat(1, encoder_input.shape[1], 1).to(module.device)
1687
 
 
120
  noised_inputs = diffusion_input * alphas + noise * sigmas
121
  targets = noise * alphas - diffusion_input * sigmas
122
 
123
+ with torch.cuda.amp.autocast():
124
  v = self.diffusion(noised_inputs, t)
125
 
126
  loss_info.update({
 
185
  noise = torch.randn([self.num_demos, module.diffusion.io_channels, demo_samples]).to(module.device)
186
 
187
  try:
188
+ with torch.cuda.amp.autocast():
189
  fakes = sample(module.diffusion_ema, noise, self.demo_steps, 0)
190
 
191
  if module.diffusion.pretransform is not None:
 
347
  self.diffusion.pretransform.to(self.device)
348
 
349
  if not self.pre_encoded:
350
+ with torch.cuda.amp.autocast() and torch.set_grad_enabled(self.diffusion.pretransform.enable_grad):
351
  self.diffusion.pretransform.train(self.diffusion.pretransform.enable_grad)
352
 
353
  diffusion_input = self.diffusion.pretransform.encode(diffusion_input)
 
467
 
468
  diffusion_input = reals
469
 
470
+ with torch.cuda.amp.autocast() and torch.no_grad():
471
  conditioning = self.diffusion.conditioner(metadata, self.device)
472
 
473
  # TODO: decide what to do with padding masks during validation
 
483
  self.diffusion.pretransform.to(self.device)
484
 
485
  if not self.pre_encoded:
486
+ with torch.cuda.amp.autocast() and torch.no_grad():
487
  self.diffusion.pretransform.train(self.diffusion.pretransform.enable_grad)
488
 
489
  diffusion_input = self.diffusion.pretransform.encode(diffusion_input)
 
522
  # if use_padding_mask:
523
  # extra_args["mask"] = padding_masks
524
 
525
+ with torch.cuda.amp.autocast() and torch.no_grad():
526
  output = self.diffusion(noised_inputs, t, cond=conditioning, cfg_dropout_prob = 0, **extra_args)
527
 
528
  val_loss = F.mse_loss(output, targets)
 
621
 
622
  try:
623
  print("Getting conditioning")
624
+ with torch.cuda.amp.autocast():
625
  conditioning = module.diffusion.conditioner(demo_cond, module.device)
626
 
627
  cond_inputs = module.diffusion.get_conditioning_inputs(conditioning)
 
665
 
666
  print(f"Generating demo for cfg scale {cfg_scale}")
667
 
668
+ with torch.cuda.amp.autocast():
669
  model = module.diffusion_ema.model if module.diffusion_ema is not None else module.diffusion.model
670
 
671
  if module.diffusion_objective == "v":
 
917
 
918
  p.tick("setup")
919
 
920
+ #with torch.cuda.amp.autocast():
921
  conditioning = self.diffusion.conditioner(metadata, self.device)
922
 
923
  p.tick("conditioning")
 
983
 
984
  extra_args = {}
985
 
986
+ #with torch.cuda.amp.autocast():
987
  p.tick("amp")
988
  output = self.diffusion(noised_inputs, t, cond=conditioning, cfg_dropout_prob = self.cfg_dropout_prob, **extra_args)
989
  p.tick("diffusion")
 
1127
  model = module.diffusion_ema.model if module.diffusion_ema is not None else module.diffusion.model
1128
  print(f"Generating demo for cfg scale {cfg_scale}")
1129
 
1130
+ with torch.cuda.amp.autocast():
1131
  if module.diffusion_objective == "v":
1132
  fakes = sample(model, noise, self.demo_steps, 0, **cond_inputs, cfg_scale=cfg_scale, dist_shift=module.diffusion.dist_shift, batch_cfg=True)
1133
  elif module.diffusion_objective == "rectified_flow":
 
1274
  noised_reals = reals * alphas + noise * sigmas
1275
  targets = noise * alphas - reals * sigmas
1276
 
1277
+ with torch.cuda.amp.autocast():
1278
  v = self.diffae.diffusion(noised_reals, t, input_concat_cond=latents)
1279
 
1280
  loss_info.update({
 
1354
 
1355
  demo_reals = demo_reals.to(module.device)
1356
 
1357
+ with torch.no_grad() and torch.cuda.amp.autocast():
1358
  latents = module.diffae_ema.ema_model.encode(encoder_input).float()
1359
  fakes = module.diffae_ema.ema_model.decode(latents, steps=self.demo_steps)
1360
 
 
1387
  audio_spectrogram_image(reals_fakes))
1388
 
1389
  if module.diffae_ema.ema_model.pretransform is not None:
1390
+ with torch.no_grad() and torch.cuda.amp.autocast():
1391
  initial_latents = module.diffae_ema.ema_model.pretransform.encode(encoder_input)
1392
  first_stage_fakes = module.diffae_ema.ema_model.pretransform.decode(initial_latents)
1393
  first_stage_fakes = rearrange(first_stage_fakes, 'b d n -> d (b n)')
 
1546
  source = self.diffusion.pretransform.encode(source)
1547
 
1548
  if self.diffusion.conditioner is not None:
1549
+ with torch.cuda.amp.autocast():
1550
  conditioning = self.diffusion.conditioner(metadata, self.device)
1551
  else:
1552
  conditioning = {}
 
1566
  noised_reals = reals * alphas + noise * sigmas
1567
  targets = noise * alphas - reals * sigmas
1568
 
1569
+ with torch.cuda.amp.autocast():
1570
 
1571
  conditioning['source'] = [source]
1572
 
 
1675
  encoder_input = demo_reals
1676
 
1677
  if module.diffusion.conditioner is not None:
1678
+ with torch.cuda.amp.autocast():
1679
  conditioning_tensors = module.diffusion.conditioner(metadata, module.device)
1680
 
1681
  else:
1682
  conditioning_tensors = {}
1683
 
1684
+ with torch.no_grad() and torch.cuda.amp.autocast():
1685
  if module.prior_type == PriorType.MonoToStereo and encoder_input.shape[1] > 1:
1686
  source = encoder_input.mean(dim=1, keepdim=True).repeat(1, encoder_input.shape[1], 1).to(module.device)
1687
 
stable-audio-tools/train.py CHANGED
@@ -145,7 +145,7 @@ def main():
145
 
146
  trainer = pl.Trainer(
147
  devices="auto",
148
- accelerator="auto",
149
  num_nodes = args.num_nodes,
150
  strategy=strategy,
151
  precision=args.precision,
 
145
 
146
  trainer = pl.Trainer(
147
  devices="auto",
148
+ accelerator="gpu",
149
  num_nodes = args.num_nodes,
150
  strategy=strategy,
151
  precision=args.precision,