cazyundee's picture
Fix sample rate to 44.1kHz and run generation on CPU (GPU decorator only on startup check and load_model)
079c016
Raw
History Blame Contribute Delete
3.34 kB
import os
import torch
import torchaudio
import gradio as gr
import spaces
from einops import rearrange
from huggingface_hub import login
from stable_audio_3 import StableAudioModel
# Authenticate with gated model when HF_TOKEN secret is present
hf_token = os.environ.get("HF_TOKEN")
if hf_token:
login(token=hf_token)
# Required for ZeroGPU Spaces - must have at least one @spaces.GPU function
@spaces.GPU(duration=1)
def _gpu_startup_check():
return "GPU check passed"
# Run the check at startup
_gpu_startup_check()
# Model cache
MODEL_CACHE = {}
@spaces.GPU(duration=1)
def load_model(model_name):
"""Load model on demand and cache it."""
if model_name not in MODEL_CACHE:
print(f"Loading {model_name} model...")
model = StableAudioModel.from_pretrained(
model_name,
device="cpu"
)
MODEL_CACHE[model_name] = model
print(f"{model_name} loaded successfully!")
return MODEL_CACHE[model_name]
def generate_audio(prompt, duration, steps, cfg_scale, seed, model_name):
print(f"Generating with {model_name}: prompt='{prompt}', duration={duration}s, steps={steps}, cfg={cfg_scale}, seed={seed}")
model = load_model(model_name)
audio = model.generate(
prompt=prompt,
duration=duration,
steps=steps,
cfg_scale=cfg_scale,
seed=seed,
batch_size=1
)
# Post-process: (batch, channels, samples) -> stereo waveform
audio = rearrange(audio, "b d n -> d (b n)")
audio = audio.to(torch.float32).clamp(-1, 1).mul(32767).to(torch.int16).cpu()
output_path = "output.wav"
torchaudio.save(output_path, audio, 44100)
print("Generation complete!")
return output_path
with gr.Blocks(title="Stable Audio 3 Small") as demo:
gr.Markdown("# 🎵 Stable Audio 3 Small - Music & SFX Generation")
gr.Markdown("Generate music and sound effects using Stability AI's Stable Audio 3 Small models. Runs on CPU.")
with gr.Row():
with gr.Column():
model_name = gr.Dropdown(
choices=["small-music", "small-sfx"],
value="small-music",
label="Model"
)
prompt = gr.Textbox(
label="Prompt",
placeholder="Describe the music or sound effect you want to generate...",
lines=2
)
duration = gr.Slider(
minimum=1, maximum=120, value=30, step=1,
label="Duration (seconds)"
)
steps = gr.Slider(
minimum=1, maximum=50, value=8, step=1,
label="Steps"
)
cfg_scale = gr.Slider(
minimum=0.0, maximum=10.0, value=1.0, step=0.1,
label="CFG Scale"
)
seed = gr.Number(
value=-1, label="Seed (-1 for random)"
)
btn = gr.Button("Generate", variant="primary")
with gr.Column():
audio_output = gr.Audio(
label="Generated Audio",
type="filepath"
)
btn.click(
fn=generate_audio,
inputs=[prompt, duration, steps, cfg_scale, seed, model_name],
outputs=audio_output
)
demo.queue(max_size=4, default_concurrency_limit=1).launch()