Actualiser server.py(gradio)
This commit is contained in:
@ -3,41 +3,65 @@ import torch
|
||||
import scipy.io.wavfile
|
||||
import numpy as np
|
||||
import tempfile
|
||||
import os
|
||||
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")
|
||||
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):
|
||||
if not prompt: return None
|
||||
music_model.set_generation_params(duration=duration)
|
||||
wav = music_model.generate([prompt])
|
||||
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))
|
||||
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))
|
||||
return path
|
||||
|
||||
def generate_sound(prompt, duration=5):
|
||||
if not prompt: return None
|
||||
audio_model.set_generation_params(duration=duration)
|
||||
wav = audio_model.generate([prompt])
|
||||
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:
|
||||
gr.Markdown("## 🎙️ Studio Audio IA")
|
||||
with gr.Tab("🎵 Musique"):
|
||||
p1 = gr.Textbox(label="Prompt")
|
||||
d1 = gr.Slider(5, 30, value=10, step=5, label="Durée (s)")
|
||||
b1 = gr.Button("Générer")
|
||||
o1 = gr.Audio(label="Résultat")
|
||||
path = os.path.join(tempfile.gettempdir(), "sound_out.wav")
|
||||
scipy.io.wavfile.write(path, audio_model.sample_rate, (wav_np * 32767).astype(np.int16))
|
||||
return path
|
||||
|
||||
# Interface UI
|
||||
with gr.Blocks(title="🎙️ Studio Audio IA", theme=gr.themes.Soft()) as demo:
|
||||
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)
|
||||
with gr.Tab("🔊 Bruitages"):
|
||||
p2 = gr.Textbox(label="Prompt")
|
||||
d2 = gr.Slider(2, 15, value=5, label="Durée (s)")
|
||||
b2 = gr.Button("Générer")
|
||||
o2 = gr.Audio(label="Résultat")
|
||||
|
||||
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)
|
||||
Reference in New Issue
Block a user