basyx commited on
Commit
7c5c205
·
verified ·
1 Parent(s): 07ddd86

Update generate.py

Browse files
Files changed (1) hide show
  1. generate.py +4 -18
generate.py CHANGED
@@ -3,26 +3,14 @@ import soundfile as sf
3
  import uuid
4
  import os
5
 
6
- from audiocraft.models import MusicGen
7
- from styles import STYLES
8
- from storage import upload_audio
9
-
10
  MODEL = None
11
 
12
 
13
  def load_model():
14
  global MODEL
15
  if MODEL is None:
 
16
  MODEL = MusicGen.get_pretrained("facebook/musicgen-small")
17
- MODEL.set_generation_params(duration=10)
18
-
19
-
20
- def build_prompt(user_prompt, style):
21
-
22
- if style and style in STYLES:
23
- return f"{STYLES[style]}, {user_prompt}"
24
-
25
- return user_prompt
26
 
27
 
28
  def generate_music(prompt, duration=10):
@@ -34,10 +22,8 @@ def generate_music(prompt, duration=10):
34
  wav = MODEL.generate([prompt])
35
 
36
  filename = f"{uuid.uuid4().hex}.wav"
37
- output_path = f"outputs/{filename}"
38
-
39
- sf.write(output_path, wav[0].cpu().numpy().T, 32000)
40
 
41
- public_url = upload_audio(output_path)
42
 
43
- return output_path, public_url
 
3
  import uuid
4
  import os
5
 
 
 
 
 
6
  MODEL = None
7
 
8
 
9
  def load_model():
10
  global MODEL
11
  if MODEL is None:
12
+ from audiocraft.models import MusicGen
13
  MODEL = MusicGen.get_pretrained("facebook/musicgen-small")
 
 
 
 
 
 
 
 
 
14
 
15
 
16
  def generate_music(prompt, duration=10):
 
22
  wav = MODEL.generate([prompt])
23
 
24
  filename = f"{uuid.uuid4().hex}.wav"
25
+ path = f"outputs/{filename}"
 
 
26
 
27
+ sf.write(path, wav[0].cpu().numpy().T, 32000)
28
 
29
+ return path