sesa-gpu / src /ui.py
John6666's picture
Upload 44 files
81ba775 verified
Raw
History Blame Contribute Delete
21.4 kB
from __future__ import annotations
import json
from pathlib import Path
import gradio as gr
from .catalog import catalog_choices, default_model_id
from .config import (
ENSEMBLE_ALGORITHMS,
JOB_TTL_SECONDS,
OUTPUT_FORMATS,
SAFE_PACKAGE_MODELS,
STEM_CHOICES,
)
from .dataset_samples import (
DEFAULT_DATASET_SAMPLE_ID,
dataset_sample_choices,
dataset_sample_markdown,
)
from .duration import (
DEFAULT_BASE_SECONDS_PER_MINUTE,
DEFAULT_MANUAL_SECONDS,
DEFAULT_SAFETY_MULTIPLIER,
DEFAULT_SAFE_MODE_ENABLED,
DURATION_MODE_MANUAL,
DURATION_MODE_SEMI_AUTO,
DURATION_MODES,
)
from .errors import StageError
from .longform import (
CHUNK_MODE_AUTO,
CHUNK_MODE_FIXED,
CHUNK_MODES,
DEFAULT_FIXED_CHUNK_SECONDS,
DEFAULT_PREVIEW_SECONDS,
RANGE_MODE_CUSTOM,
RANGE_MODE_FULL,
RANGE_MODE_PREVIEW,
RANGE_MODES,
)
from .observability import DEFAULT_LOG_LEVEL, LOG_LEVELS
from .preparation import (
BITRATE_AUTO,
BUNDLE_LAYOUT_BY_INPUT,
BUNDLE_LAYOUT_FLAT,
cleanup_prepared_state,
prepare_job,
)
from .service import (
EXECUTION_MODE_INITIAL,
EXECUTION_MODE_RETRY_INCOMPLETE,
request_prepared_job_cancel,
run_prepared_job,
)
OUTPUT_BITRATES = [BITRATE_AUTO, "96k", "128k", "192k", "256k", "320k"]
OUTPUT_SAMPLE_RATES = [24000, 32000, 44100, 48000]
def _result_has_retryable_items(result) -> bool:
try:
returned_files = result[3] or []
except Exception:
return False
for value in returned_files:
path = Path(str(value))
if path.name != "batch_manifest.json" or not path.is_file():
continue
try:
payload = json.loads(path.read_text(encoding="utf-8"))
except Exception:
return False
return bool(payload.get("retryable_input_indexes"))
return False
def run_prepared_separation(
prepared_state_value,
execution_mode=EXECUTION_MODE_INITIAL,
progress=None,
callback_timing=None,
):
result = run_prepared_job(
prepared_state_value,
progress=progress,
callback_timing=callback_timing,
execution_mode=execution_mode,
)
retryable = _result_has_retryable_items(result)
return (
*result,
gr.update(interactive=False),
gr.update(interactive=retryable),
gr.update(interactive=retryable),
)
def request_cancel_with_ui(prepared_state_value):
return request_prepared_job_cancel(prepared_state_value)
def prepare_separation_with_progress(*args, progress=gr.Progress(track_tqdm=True)):
try:
if len(args) not in {34, 35}:
raise ValueError(f"Expected 34 or 35 Preparation inputs, received {len(args)}.")
batch_continue_on_item_error = args[34] if len(args) == 35 else True
state, markdown, config_path, log_path = prepare_job(
*args[:22],
0,
*args[28:34],
progress=progress,
range_mode=args[22],
range_start_seconds=args[23],
range_end_seconds=args[24],
preview_seconds=args[25],
chunk_mode=args[26],
fixed_chunk_seconds=args[27],
batch_continue_on_item_error=batch_continue_on_item_error,
)
return (
state, markdown, config_path, log_path,
gr.update(interactive=True), gr.update(interactive=False), gr.update(interactive=True),
)
except StageError as exc:
hint = f"\n\n{exc.hint}" if exc.hint else ""
diagnostic = getattr(exc, "diagnostic_path", None)
log_path = getattr(exc, "log_path", None)
markdown = (
"### Preparation failed before GPU allocation\n"
f"- Stage: `{exc.stage}`\n"
f"- Error: `{exc.message}`{hint}"
)
return (
"", markdown, diagnostic, log_path,
gr.update(interactive=False), gr.update(interactive=False), gr.update(interactive=False),
)
except Exception as exc:
return (
"",
f"### Preparation failed before GPU allocation\n`{type(exc).__name__}: {exc}`",
None,
None,
gr.update(interactive=False),
gr.update(interactive=False),
gr.update(interactive=False),
)
def _source_visibility(value: str):
return (
gr.update(visible=value == "Hugging Face"),
gr.update(visible=value == "GitHub"),
)
def _range_visibility(value: str):
return (
gr.update(visible=value != RANGE_MODE_FULL),
gr.update(visible=value == RANGE_MODE_CUSTOM),
gr.update(visible=value == RANGE_MODE_PREVIEW),
)
def _chunk_visibility(value: str):
return gr.update(visible=value == CHUNK_MODE_FIXED)
def _duration_visibility(value: str):
return (
gr.update(visible=value == DURATION_MODE_MANUAL),
gr.update(visible=value != DURATION_MODE_MANUAL),
)
def _duration_controls(label_prefix: str = ""):
prefix = f"{label_prefix} " if label_prefix else ""
duration_mode = gr.Radio(
label=f"{prefix}Time request mode",
choices=DURATION_MODES,
value=DURATION_MODE_SEMI_AUTO,
)
with gr.Group(visible=False) as manual_group:
manual_gpu_seconds = gr.Slider(
label=f"{prefix}Manual maximum GPU time (seconds)",
minimum=30,
maximum=300,
step=5,
value=DEFAULT_MANUAL_SECONDS,
)
with gr.Group(visible=True) as semi_auto_group:
semi_auto_safe_mode = gr.Checkbox(
label=f"{prefix}Extra safety margin (+30%, Semi-auto only)",
value=DEFAULT_SAFE_MODE_ENABLED,
)
gr.Markdown(
"The extra margin defaults OFF. A request that is too short may still finish, but it can be aborted after an overrun warning. "
"The +30% margin affects only Semi-auto; Manual is unchanged."
)
with gr.Accordion("Advanced Semi-auto calibration", open=False):
semi_auto_base_seconds = gr.Number(
label=f"{prefix}Base estimate: GPU seconds per source minute and model",
value=DEFAULT_BASE_SECONDS_PER_MINUTE,
minimum=1,
maximum=120,
precision=1,
)
semi_auto_safety_multiplier = gr.Slider(
label=f"{prefix}Calibration multiplier",
minimum=0.5,
maximum=3.0,
step=0.05,
value=DEFAULT_SAFETY_MULTIPLIER,
)
duration_mode.change(
_duration_visibility,
inputs=duration_mode,
outputs=[manual_group, semi_auto_group],
show_progress="hidden",
)
return (
duration_mode,
manual_gpu_seconds,
semi_auto_base_seconds,
semi_auto_safety_multiplier,
semi_auto_safe_mode,
)
def _invalidate_preparation():
return (
"",
"Settings changed. Run **1. Prepare / refresh plan** again before requesting GPU time.",
gr.update(interactive=False),
gr.update(interactive=False),
gr.update(interactive=False),
)
def build_demo(
submit_function,
prepare_function=prepare_separation_with_progress,
):
choices = catalog_choices()
sample_choices = [("Use upload only", ""), *dataset_sample_choices()]
with gr.Blocks(title="SESA Fast Separation", delete_cache=(3600, 21600)) as demo:
gr.Markdown(
"# SESA Fast Separation\n"
"Upload one or more audio or video files, prepare a frozen job plan, then run separation on GPU. "
"Preparation validates media, extracts the selected range, downloads models, and creates an isolated job before GPU time is requested."
)
with gr.Row():
with gr.Column(scale=3):
files = gr.File(
label="Audio or video files",
file_count="multiple",
type="filepath",
)
public_sample = gr.Dropdown(
label="Or use a public sample",
choices=sample_choices,
value="",
interactive=True,
)
public_sample_info = gr.Markdown(dataset_sample_markdown(""))
public_sample.change(
dataset_sample_markdown,
inputs=public_sample,
outputs=public_sample_info,
show_progress="hidden",
)
gr.Markdown("Select up to 6 models total across curated, built-in, and custom sources.")
catalog_models = gr.Dropdown(
label="Curated models",
choices=choices,
value=[default_model_id()],
multiselect=True,
max_choices=6,
filterable=True,
)
package_models = gr.Dropdown(
label="Additional built-in models",
choices=SAFE_PACKAGE_MODELS,
multiselect=True,
max_choices=6,
value=[],
)
with gr.Column(scale=2):
ensemble_algorithm = gr.Dropdown(
label="Ensemble algorithm",
choices=ENSEMBLE_ALGORITHMS,
value="avg_wave",
)
output_format = gr.Dropdown(
label="Output format", choices=OUTPUT_FORMATS, value="FLAC"
)
output_bitrate = gr.Dropdown(
label="Output bitrate (compressed formats)",
choices=OUTPUT_BITRATES,
value=BITRATE_AUTO,
)
output_sample_rate = gr.Dropdown(
label="Output sample rate",
choices=OUTPUT_SAMPLE_RATES,
value=44100,
)
normalization_threshold = gr.Slider(
label="Normalization peak threshold",
minimum=0.1,
maximum=1.0,
step=0.01,
value=0.9,
)
amplification_threshold = gr.Slider(
label="Amplification minimum peak threshold",
minimum=0.0,
maximum=1.0,
step=0.01,
value=0.0,
)
single_stem = gr.Dropdown(
label="Output stems", choices=STEM_CHOICES, value="All stems"
)
bundle_layout = gr.Radio(
label="Result ZIP layout",
choices=[BUNDLE_LAYOUT_FLAT, BUNDLE_LAYOUT_BY_INPUT],
value=BUNDLE_LAYOUT_FLAT,
)
allow_cpu_fallback = gr.Checkbox(
label="Allow limited CPU fallback when CUDA is unavailable",
value=False,
)
pitch_shift = gr.Slider(
label="MDXC pitch shift (semitones)", minimum=-12, maximum=12, step=1, value=0
)
with gr.Accordion("Range and long-file chunking", open=False):
gr.Markdown(
"Range extraction runs during Preparation, before ZeroGPU. The selected chunk plan is resolved and saved in the frozen job plan."
)
range_mode = gr.Radio(
label="Processing range", choices=RANGE_MODES, value=RANGE_MODE_FULL
)
with gr.Group(visible=False) as range_start_group:
range_start_seconds = gr.Number(
label="Start position (seconds)", value=0, minimum=0, precision=3
)
with gr.Group(visible=False) as range_end_group:
range_end_seconds = gr.Number(
label="End position (seconds; 0 means source end)", value=0, minimum=0, precision=3
)
with gr.Group(visible=False) as preview_group:
preview_seconds = gr.Slider(
label="Preview range length (seconds)", minimum=5, maximum=600, step=5, value=DEFAULT_PREVIEW_SECONDS
)
range_mode.change(
_range_visibility,
inputs=range_mode,
outputs=[range_start_group, range_end_group, preview_group],
show_progress="hidden",
)
chunk_mode = gr.Radio(
label="Long-file chunk policy", choices=CHUNK_MODES, value=CHUNK_MODE_AUTO
)
with gr.Group(visible=False) as fixed_chunk_group:
fixed_chunk_seconds = gr.Dropdown(
label="Fixed chunk duration",
choices=[("2 minutes", 120), ("5 minutes", 300), ("10 minutes", 600)],
value=DEFAULT_FIXED_CHUNK_SECONDS,
)
chunk_mode.change(
_chunk_visibility,
inputs=chunk_mode,
outputs=fixed_chunk_group,
show_progress="hidden",
)
gr.Markdown(
"Chunk results are concatenated by audio-separator without crossfade; rare boundary artifacts are possible."
)
with gr.Accordion("Batch behavior", open=False):
gr.Markdown(
"Upload order is preserved. Models load once per prepared job and inputs run sequentially in the same GPU callback. "
"The Batch manifest is returned with the outputs and is also included in the result ZIP. "
"Stop is cooperative at item boundaries; failed or unprocessed items can be retried from the same frozen Preparation."
)
batch_continue_on_item_error = gr.Checkbox(
label="Continue with remaining files if one file fails",
value=True,
)
with gr.Accordion("Custom MDXC / Roformer model", open=False):
custom_source = gr.Radio(
label="Source", choices=["None", "Hugging Face", "GitHub"], value="None"
)
with gr.Group(visible=False) as hf_group:
hf_repo = gr.Textbox(label="Hugging Face repository", placeholder="owner/repository")
hf_weight = gr.Textbox(label="Weight filename", placeholder="model.ckpt")
hf_config = gr.Textbox(label="YAML filename", placeholder="config.yaml")
hf_revision = gr.Textbox(label="Revision", value="main")
with gr.Group(visible=False) as github_group:
github_weight_url = gr.Textbox(
label="GitHub release asset URL",
placeholder="https://github.com/owner/repo/releases/download/tag/model.ckpt",
)
github_config_url = gr.Textbox(
label="GitHub YAML URL",
placeholder="https://raw.githubusercontent.com/owner/repo/commit/config.yaml",
)
github_sha256 = gr.Textbox(label="Weight SHA-256 (recommended)")
custom_source.change(
_source_visibility,
inputs=custom_source,
outputs=[hf_group, github_group],
show_progress="hidden",
)
with gr.Accordion("ZeroGPU time request", open=False):
(
duration_mode,
manual_gpu_seconds,
semi_auto_base_seconds,
semi_auto_safety_multiplier,
semi_auto_safe_mode,
) = _duration_controls()
with gr.Accordion("Progress and logging", open=False):
log_level = gr.Radio(
label="Log detail", choices=LOG_LEVELS, value=DEFAULT_LOG_LEVEL
)
gr.Markdown(
"Preparation and GPU execution use separate per-job logs. INFO is compact; DEBUG includes additional package progress detail."
)
prepare_inputs = [
files,
public_sample,
catalog_models,
package_models,
custom_source,
hf_repo,
hf_weight,
hf_config,
hf_revision,
github_weight_url,
github_config_url,
github_sha256,
ensemble_algorithm,
output_format,
output_bitrate,
output_sample_rate,
normalization_threshold,
amplification_threshold,
single_stem,
bundle_layout,
allow_cpu_fallback,
pitch_shift,
range_mode,
range_start_seconds,
range_end_seconds,
preview_seconds,
chunk_mode,
fixed_chunk_seconds,
duration_mode,
manual_gpu_seconds,
semi_auto_base_seconds,
semi_auto_safety_multiplier,
semi_auto_safe_mode,
log_level,
batch_continue_on_item_error,
]
prepared_state = gr.State(
"",
time_to_live=JOB_TTL_SECONDS,
delete_callback=cleanup_prepared_state,
)
with gr.Row():
prepare_button = gr.Button("1. Prepare / refresh plan", variant="secondary")
run_button = gr.Button("2. Separate prepared job", variant="primary", interactive=False)
with gr.Row():
retry_button = gr.Button("Retry incomplete Batch items", variant="secondary", interactive=False)
stop_button = gr.Button("Stop after current Batch item", variant="stop", interactive=False)
initial_mode = gr.State(EXECUTION_MODE_INITIAL)
retry_mode = gr.State(EXECUTION_MODE_RETRY_INCOMPLETE)
cancel_status = gr.Markdown()
preflight = gr.Markdown("Prepare the job before requesting GPU time.")
with gr.Row():
preflight_json = gr.File(label="Prepared job plan (JSON)")
preparation_log = gr.File(label="Preparation log")
prepare_button.click(
prepare_function,
inputs=prepare_inputs,
outputs=[
prepared_state, preflight, preflight_json, preparation_log,
run_button, retry_button, stop_button,
],
api_name="prepare",
concurrency_limit=2,
concurrency_id="sesa_prepare",
show_progress="full",
)
# Any setting change invalidates the frozen plan. The existing job remains
# isolated until TTL cleanup, but cannot be submitted through this session state.
for component in prepare_inputs:
component.change(
_invalidate_preparation,
inputs=None,
outputs=[prepared_state, preflight, run_button, retry_button, stop_button],
show_progress="hidden",
)
status = gr.Markdown()
with gr.Row():
preview_one = gr.Audio(label="Output preview 1", type="filepath")
preview_two = gr.Audio(label="Output preview 2", type="filepath")
output_files = gr.File(label="Output files and Batch manifest", file_count="multiple")
output_zip = gr.File(label="Result ZIP")
with gr.Row():
reproducibility_json = gr.File(label="Reproducibility record (JSON)")
job_log_file = gr.File(label="Runtime log file")
with gr.Accordion("Runtime log details", open=False):
job_log_tail = gr.Textbox(
label="Runtime log tail", lines=14, max_lines=28, interactive=False, autoscroll=True
)
execution_outputs = [
status,
preview_one,
preview_two,
output_files,
output_zip,
job_log_tail,
job_log_file,
reproducibility_json,
run_button,
retry_button,
stop_button,
]
run_event = run_button.click(
submit_function,
inputs=[prepared_state, initial_mode],
outputs=execution_outputs,
api_name="separate_prepared",
concurrency_limit=1,
concurrency_id="sesa_gpu",
show_progress="full",
)
retry_event = retry_button.click(
submit_function,
inputs=[prepared_state, retry_mode],
outputs=execution_outputs,
api_name="retry_prepared_batch",
concurrency_limit=1,
concurrency_id="sesa_gpu",
show_progress="full",
)
stop_button.click(
request_cancel_with_ui,
inputs=[prepared_state],
outputs=[cancel_status],
cancels=[run_event, retry_event],
queue=False,
api_name="request_batch_cancel",
api_visibility="private",
show_progress="hidden",
)
gr.Markdown(
"Only the model cache is shared. Inputs, outputs, logs, and job configuration are isolated per job and removed automatically after the retention period."
)
return demo