67 lines
2.9 KiB
Plaintext
67 lines
2.9 KiB
Plaintext
import gradio as gr
|
|
import torch
|
|
import scipy.io.wavfile
|
|
import numpy as np
|
|
import tempfile
|
|
import os
|
|
from audiocraft.models import MusicGen, AudioGen
|
|
|
|
# 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()
|
|
|
|
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()
|
|
|
|
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 (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) |