basyx commited on
Commit
83f27a4
·
verified ·
1 Parent(s): d7c04ff

Update generate.py

Browse files
Files changed (1) hide show
  1. generate.py +25 -15
generate.py CHANGED
@@ -1,29 +1,39 @@
1
  import torch
 
2
  import soundfile as sf
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):
17
 
18
- load_model()
 
 
 
 
19
 
20
- MODEL.set_generation_params(duration=duration)
 
 
 
21
 
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
 
1
  import torch
2
+ from transformers import AutoProcessor, MusicgenForConditionalGeneration
3
  import soundfile as sf
4
  import uuid
 
5
 
6
+ MODEL_NAME = "facebook/musicgen-small"
7
 
8
+ print("Loading MusicGen model...")
9
 
10
+ processor = AutoProcessor.from_pretrained(MODEL_NAME)
11
+ model = MusicgenForConditionalGeneration.from_pretrained(MODEL_NAME)
 
 
 
12
 
13
+ device = "cuda" if torch.cuda.is_available() else "cpu"
14
+ model.to(device)
15
 
16
+ def generate_music(prompt, duration):
17
 
18
+ inputs = processor(
19
+ text=[prompt],
20
+ padding=True,
21
+ return_tensors="pt"
22
+ ).to(device)
23
 
24
+ audio_values = model.generate(
25
+ **inputs,
26
+ max_new_tokens=int(duration * 50)
27
+ )
28
 
29
+ filename = f"/tmp/{uuid.uuid4()}.wav"
30
 
31
+ sampling_rate = model.config.audio_encoder.sampling_rate
 
32
 
33
+ sf.write(
34
+ filename,
35
+ audio_values[0, 0].cpu().numpy(),
36
+ sampling_rate
37
+ )
38
 
39
+ return filename