basyx commited on
Commit
fef7c67
·
verified ·
1 Parent(s): 8c2aa75

Update generate.py

Browse files
Files changed (1) hide show
  1. generate.py +25 -30
generate.py CHANGED
@@ -1,44 +1,39 @@
1
  import torch
2
- from audiocraft.models import MusicGen
3
- from audiocraft.data.audio import audio_write
4
- from pydub import AudioSegment
5
- import uuid, os
6
 
7
- DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
8
 
9
- print("Loading MusicGen...")
10
- model = MusicGen.get_pretrained("facebook/musicgen-small")
11
 
12
- OUTPUT_DIR = "outputs"
13
- os.makedirs(OUTPUT_DIR, exist_ok=True)
14
 
 
 
15
 
16
  def generate_music(prompt, duration):
17
 
18
- model.set_generation_params(duration=duration)
 
 
 
 
19
 
20
- wav = model.generate([prompt])
21
-
22
- name = str(uuid.uuid4())
23
- wav_path = f"{OUTPUT_DIR}/{name}"
24
-
25
- audio_write(
26
- wav_path,
27
- wav[0].cpu(),
28
- model.sample_rate,
29
- strategy="loudness",
30
- loudness_compressor=True,
31
  )
32
 
33
- # Convert to MP3 (smaller + faster)
34
- mp3_path = f"{OUTPUT_DIR}/{name}.mp3"
35
 
36
- AudioSegment.from_wav(f"{wav_path}.wav").export(
37
- mp3_path,
38
- format="mp3",
39
- bitrate="192k"
40
- )
41
 
42
- os.remove(f"{wav_path}.wav")
 
 
 
 
43
 
44
- return mp3_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