Actualiser server.py(gradio)
This commit is contained in:
@ -3,41 +3,65 @@ import torch
|
|||||||
import scipy.io.wavfile
|
import scipy.io.wavfile
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import tempfile
|
import tempfile
|
||||||
|
import os
|
||||||
from audiocraft.models import MusicGen, AudioGen
|
from audiocraft.models import MusicGen, AudioGen
|
||||||
|
|
||||||
print("🔄 Chargement des modèles...")
|
# Force l'utilisation du GPU si disponible
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
|
||||||
|
print(f"🔄 Chargement des modèles sur {device}...")
|
||||||
|
try:
|
||||||
|
# On utilise 'small' pour la musique et 'medium' pour l'audio pour un bon ratio vitesse/qualité
|
||||||
music_model = MusicGen.get_pretrained("facebook/musicgen-small")
|
music_model = MusicGen.get_pretrained("facebook/musicgen-small")
|
||||||
audio_model = AudioGen.get_pretrained("facebook/audiogen-medium")
|
audio_model = AudioGen.get_pretrained("facebook/audiogen-medium")
|
||||||
|
except Exception as e:
|
||||||
|
print(f"❌ Erreur lors du chargement des modèles : {e}")
|
||||||
|
|
||||||
def generate_music(prompt, duration=10):
|
def generate_music(prompt, duration=10):
|
||||||
|
if not prompt: return None
|
||||||
music_model.set_generation_params(duration=duration)
|
music_model.set_generation_params(duration=duration)
|
||||||
wav = music_model.generate([prompt])
|
wav = music_model.generate([prompt])
|
||||||
wav_np = wav[0, 0].cpu().numpy()
|
wav_np = wav[0, 0].cpu().numpy()
|
||||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
|
|
||||||
scipy.io.wavfile.write(f.name, music_model.sample_rate, (wav_np * 32767).astype(np.int16))
|
path = os.path.join(tempfile.gettempdir(), "music_out.wav")
|
||||||
return f.name
|
scipy.io.wavfile.write(path, music_model.sample_rate, (wav_np * 32767).astype(np.int16))
|
||||||
|
return path
|
||||||
|
|
||||||
def generate_sound(prompt, duration=5):
|
def generate_sound(prompt, duration=5):
|
||||||
|
if not prompt: return None
|
||||||
audio_model.set_generation_params(duration=duration)
|
audio_model.set_generation_params(duration=duration)
|
||||||
wav = audio_model.generate([prompt])
|
wav = audio_model.generate([prompt])
|
||||||
wav_np = wav[0, 0].cpu().numpy()
|
wav_np = wav[0, 0].cpu().numpy()
|
||||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
|
|
||||||
scipy.io.wavfile.write(f.name, audio_model.sample_rate, (wav_np * 32767).astype(np.int16))
|
|
||||||
return f.name
|
|
||||||
|
|
||||||
with gr.Blocks(title="🎙️ Studio Audio IA") as demo:
|
path = os.path.join(tempfile.gettempdir(), "sound_out.wav")
|
||||||
gr.Markdown("## 🎙️ Studio Audio IA")
|
scipy.io.wavfile.write(path, audio_model.sample_rate, (wav_np * 32767).astype(np.int16))
|
||||||
with gr.Tab("🎵 Musique"):
|
return path
|
||||||
p1 = gr.Textbox(label="Prompt")
|
|
||||||
d1 = gr.Slider(5, 30, value=10, step=5, label="Durée (s)")
|
# Interface UI
|
||||||
b1 = gr.Button("Générer")
|
with gr.Blocks(title="🎙️ Studio Audio IA", theme=gr.themes.Soft()) as demo:
|
||||||
o1 = gr.Audio(label="Résultat")
|
gr.Markdown("# 🎙️ Studio Audio IA")
|
||||||
|
gr.Markdown("Générez de la musique et des effets sonores avec les modèles de Facebook Research.")
|
||||||
|
|
||||||
|
with gr.Tab("🎵 Musique (MusicGen)"):
|
||||||
|
with gr.Row():
|
||||||
|
with gr.Column():
|
||||||
|
p1 = gr.Textbox(label="Description de la musique", placeholder="Ex: Epic cinematic orchestral music, high strings...")
|
||||||
|
d1 = gr.Slider(minimum=1, maximum=30, value=10, step=1, label="Durée (secondes)")
|
||||||
|
b1 = gr.Button("Générer la musique", variant="primary")
|
||||||
|
with gr.Column():
|
||||||
|
o1 = gr.Audio(label="Résultat Audio", type="filepath")
|
||||||
b1.click(generate_music, inputs=[p1, d1], outputs=o1)
|
b1.click(generate_music, inputs=[p1, d1], outputs=o1)
|
||||||
with gr.Tab("🔊 Bruitages"):
|
|
||||||
p2 = gr.Textbox(label="Prompt")
|
with gr.Tab("🔊 Bruitages (AudioGen)"):
|
||||||
d2 = gr.Slider(2, 15, value=5, label="Durée (s)")
|
with gr.Row():
|
||||||
b2 = gr.Button("Générer")
|
with gr.Column():
|
||||||
o2 = gr.Audio(label="Résultat")
|
p2 = gr.Textbox(label="Effet sonore", placeholder="Ex: Dogs barking, car engine revving, rain on window...")
|
||||||
|
d2 = gr.Slider(minimum=1, maximum=15, value=5, step=1, label="Durée (secondes)")
|
||||||
|
b2 = gr.Button("Générer le son", variant="primary")
|
||||||
|
with gr.Column():
|
||||||
|
o2 = gr.Audio(label="Résultat Audio", type="filepath")
|
||||||
b2.click(generate_sound, inputs=[p2, d2], outputs=o2)
|
b2.click(generate_sound, inputs=[p2, d2], outputs=o2)
|
||||||
|
|
||||||
|
# Important : server_name="0.0.0.0" permet l'accès externe depuis Vast.ai
|
||||||
|
if __name__ == "__main__":
|
||||||
demo.launch(server_name="0.0.0.0", server_port=7860, share=False)
|
demo.launch(server_name="0.0.0.0", server_port=7860, share=False)
|
||||||
Reference in New Issue
Block a user