SAORI / app.py
alfredplpl's picture
Update app.py
a853813 verified
Raw
History Blame Contribute Delete
4.25 kB
import spaces
import os
import random
import gradio as gr
import numpy as np
import torch
from diffusers import StableAudioPipeline
MODEL_ID = "alfredplpl/SAO-Re-2"
token=os.environ["HF_TOKEN"]
device = "cuda"
dtype = torch.bfloat16
pipe = StableAudioPipeline.from_pretrained(
MODEL_ID,
torch_dtype=dtype,
token=token
)
pipe = pipe.to(device)
if device == "cuda":
pipe.enable_model_cpu_offload()
@spaces.GPU
def generate_audio(
prompt: str,
negative_prompt: str,
duration: float,
steps: int,
guidance_scale: float,
seed: int,
):
if not prompt.strip():
raise gr.Error("Promptを入力してください。")
if seed < 0:
seed = random.randint(0, 2**32 - 1)
generator = torch.Generator(device=device).manual_seed(seed)
with torch.inference_mode():
audio = pipe(
prompt=prompt,
negative_prompt=negative_prompt or None,
num_inference_steps=steps,
guidance_scale=guidance_scale,
audio_end_in_s=duration,
num_waveforms_per_prompt=1,
generator=generator,
).audios[0]
# diffusers StableAudio: [channels, samples] torch tensor の想定
audio = audio.T.float().cpu().numpy()
# Gradio向けに軽く安全化
audio = np.nan_to_num(audio)
peak = np.max(np.abs(audio))
if peak > 1.0:
audio = audio / peak
sample_rate = pipe.vae.sampling_rate
return (sample_rate, audio), seed
with gr.Blocks(title="Text-to-Audio") as demo:
gr.Markdown(
"""
# Text-to-Audio Demo
Text-to-Audio Demo of Stable Audio Open ReImplementation. You may use this output for commercial purposes.
"""
)
with gr.Row():
with gr.Column():
prompt = gr.Textbox(
label="Prompt",
value="A woman is speaking.",
lines=2,
)
negative_prompt = gr.Textbox(
label="Negative Prompt",
value="low quality, quality, music",
lines=2,
)
duration = gr.Slider(
minimum=1.0,
maximum=10.0,
value=5.0,
step=0.5,
label="Duration seconds",
)
steps = gr.Slider(
minimum=10,
maximum=300,
value=200,
step=1,
label="Inference steps",
)
guidance_scale = gr.Slider(
minimum=1.0,
maximum=15.0,
value=7.0,
step=0.5,
label="Guidance scale",
)
seed = gr.Number(
label="Seed -1 means random",
value=-1,
precision=0,
)
button = gr.Button("Generate", variant="primary")
with gr.Column():
output_audio = gr.Audio(
label="Generated Audio",
type="numpy",
)
used_seed = gr.Number(
label="Used Seed",
precision=0,
)
button.click(
fn=generate_audio,
inputs=[
prompt,
negative_prompt,
duration,
steps,
guidance_scale,
seed,
],
outputs=[
output_audio,
used_seed,
],
)
gr.Examples(
examples=[
[
"Playing a guitar.",
"low quality, noisy",
5.0,
200,
7.0,
-1,
],
[
"Playing a piano.",
"low quality, noisy",
5.0,
200,
7.0,
-1,
],
[
"wind",
"",
5.0,
200,
7.0,
-1,
],
],
inputs=[
prompt,
negative_prompt,
duration,
steps,
guidance_scale,
seed,
],
)
if __name__ == "__main__":
demo.queue().launch()