Actualiser server.py(gradio)

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

View File

@ -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)