Actualiser server.py(gradio)

This commit is contained in:
2026-02-23 13:27:46 +00:00
parent 2c225b436a
commit 624b5a150b

View File

@ -1,43 +1,67 @@
import gradio as gr import gradio as gr
import torch import torch
import scipy.io.wavfile import scipy.io.wavfile
import numpy as np import numpy as np
import tempfile import tempfile
from audiocraft.models import MusicGen, AudioGen import os
from audiocraft.models import MusicGen, AudioGen
print("🔄 Chargement des modèles...")
music_model = MusicGen.get_pretrained("facebook/musicgen-small") # Force l'utilisation du GPU si disponible
audio_model = AudioGen.get_pretrained("facebook/audiogen-medium") device = "cuda" if torch.cuda.is_available() else "cpu"
def generate_music(prompt, duration=10): print(f"🔄 Chargement des modèles sur {device}...")
music_model.set_generation_params(duration=duration) try:
wav = music_model.generate([prompt]) # On utilise 'small' pour la musique et 'medium' pour l'audio pour un bon ratio vitesse/qualité
wav_np = wav[0, 0].cpu().numpy() music_model = MusicGen.get_pretrained("facebook/musicgen-small")
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f: audio_model = AudioGen.get_pretrained("facebook/audiogen-medium")
scipy.io.wavfile.write(f.name, music_model.sample_rate, (wav_np * 32767).astype(np.int16)) except Exception as e:
return f.name print(f"❌ Erreur lors du chargement des modèles : {e}")
def generate_sound(prompt, duration=5): def generate_music(prompt, duration=10):
audio_model.set_generation_params(duration=duration) if not prompt: return None
wav = audio_model.generate([prompt]) music_model.set_generation_params(duration=duration)
wav_np = wav[0, 0].cpu().numpy() wav = music_model.generate([prompt])
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f: wav_np = wav[0, 0].cpu().numpy()
scipy.io.wavfile.write(f.name, audio_model.sample_rate, (wav_np * 32767).astype(np.int16))
return f.name path = os.path.join(tempfile.gettempdir(), "music_out.wav")
scipy.io.wavfile.write(path, music_model.sample_rate, (wav_np * 32767).astype(np.int16))
with gr.Blocks(title="🎙️ Studio Audio IA") as demo: return path
gr.Markdown("## 🎙️ Studio Audio IA")
with gr.Tab("🎵 Musique"): def generate_sound(prompt, duration=5):
p1 = gr.Textbox(label="Prompt") if not prompt: return None
d1 = gr.Slider(5, 30, value=10, step=5, label="Durée (s)") audio_model.set_generation_params(duration=duration)
b1 = gr.Button("Générer") wav = audio_model.generate([prompt])
o1 = gr.Audio(label="Résultat") wav_np = wav[0, 0].cpu().numpy()
b1.click(generate_music, inputs=[p1, d1], outputs=o1)
with gr.Tab("🔊 Bruitages"): path = os.path.join(tempfile.gettempdir(), "sound_out.wav")
p2 = gr.Textbox(label="Prompt") scipy.io.wavfile.write(path, audio_model.sample_rate, (wav_np * 32767).astype(np.int16))
d2 = gr.Slider(2, 15, value=5, label="Durée (s)") return path
b2 = gr.Button("Générer")
o2 = gr.Audio(label="Résultat") # Interface UI
b2.click(generate_sound, inputs=[p2, d2], outputs=o2) with gr.Blocks(title="🎙️ Studio Audio IA", theme=gr.themes.Soft()) as demo:
gr.Markdown("# 🎙️ Studio Audio IA")
demo.launch(server_name="0.0.0.0", server_port=7860, share=False) 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)
with gr.Tab("🔊 Bruitages (AudioGen)"):
with gr.Row():
with gr.Column():
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)
# 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)