cazyundee commited on
Commit
9286e36
·
verified ·
1 Parent(s): 7402b0c

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +24 -3
app.py CHANGED
@@ -2,15 +2,36 @@ import os
2
  import torch
3
  import torchaudio
4
  import gradio as gr
 
5
  from einops import rearrange
6
  from stable_audio_3 import StableAudioModel
7
 
8
- # Load model once at startup
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
  print("Loading Stable Audio 3 Small model...")
10
- model = StableAudioModel.from_pretrained("small-music", device="cpu")
 
 
 
 
11
  print("Model loaded successfully!")
12
 
13
 
 
14
  def generate_audio(prompt, duration, steps, cfg_scale, seed):
15
  print(f"Generating: prompt='{prompt}', duration={duration}s, steps={steps}, cfg={cfg_scale}, seed={seed}")
16
 
@@ -28,7 +49,7 @@ def generate_audio(prompt, duration, steps, cfg_scale, seed):
28
  audio = audio.to(torch.float32).clamp(-1, 1).mul(32767).to(torch.int16).cpu()
29
 
30
  output_path = "output.wav"
31
- torchaudio.save(output_path, audio, 44100)
32
  print("Generation complete!")
33
  return output_path
34
 
 
2
  import torch
3
  import torchaudio
4
  import gradio as gr
5
+ import spaces
6
  from einops import rearrange
7
  from stable_audio_3 import StableAudioModel
8
 
9
+ # Get token from HF Spaces secrets
10
+ hf_token = os.environ.get("HF_TOKEN")
11
+
12
+
13
+ # Required for ZeroGPU Spaces - must have at least one @spaces.GPU function
14
+ # This is a minimal startup check that uses negligible quota
15
+ @spaces.GPU(duration=1)
16
+ def _gpu_startup_check():
17
+ return "GPU check passed"
18
+
19
+
20
+ # Run the check at startup
21
+ _gpu_startup_check()
22
+
23
+
24
+ # Load model once at startup (on CPU)
25
  print("Loading Stable Audio 3 Small model...")
26
+ model = StableAudioModel.from_pretrained(
27
+ "small-music",
28
+ device="cpu",
29
+ token=hf_token
30
+ )
31
  print("Model loaded successfully!")
32
 
33
 
34
+ @spaces.GPU(duration=1)
35
  def generate_audio(prompt, duration, steps, cfg_scale, seed):
36
  print(f"Generating: prompt='{prompt}', duration={duration}s, steps={steps}, cfg={cfg_scale}, seed={seed}")
37
 
 
49
  audio = audio.to(torch.float32).clamp(-1, 1).mul(32767).to(torch.int16).cpu()
50
 
51
  output_path = "output.wav"
52
+ torchaudio.save(output_path, audio, 44000)
53
  print("Generation complete!")
54
  return output_path
55