Spaces:
Paused
Paused
| import gradio as gr | |
| import spaces | |
| import tempfile | |
| import numpy as np | |
| import scipy.io.wavfile | |
| import torch | |
| print("Cargando modelo...") | |
| from transformers import AutoProcessor, MusicgenForConditionalGeneration | |
| processor = AutoProcessor.from_pretrained("facebook/musicgen-medium") | |
| modelo = MusicgenForConditionalGeneration.from_pretrained("facebook/musicgen-medium") | |
| SR = modelo.config.audio_encoder.sampling_rate # 32000 | |
| print("Modelo listo") | |
| FRAME = SR // 50 # muestras por token (musicgen va a 50 Hz) | |
| MAX_TOKENS = 1500 # tope fisico de musicgen (~30s) | |
| CONTEXT_SEG = 5 # segundos de "arranque" que damos para continuar | |
| CHUNK_MAX_SEG = 25 # maximo que puede generar por pasada con ese contexto | |
| def _a_numpy(audio_values): | |
| a = audio_values[0].cpu().numpy() | |
| return np.squeeze(a).astype(np.float32) | |
| def _crossfade_join(a, b, sr, xf_seg=0.4): | |
| """Une a + b con un cruce corto para tapar la costura.""" | |
| xf = min(int(xf_seg * sr), len(a), len(b)) | |
| if xf <= 0: | |
| return np.concatenate([a, b]) | |
| fade_out = np.linspace(1.0, 0.0, xf) | |
| fade_in = np.linspace(0.0, 1.0, xf) | |
| mezcla = a[-xf:] * fade_out + b[:xf] * fade_in | |
| return np.concatenate([a[:-xf], mezcla, b[xf:]]) | |
| def generar_musica(prompt: str, duracion: int = 30) -> str: | |
| if not prompt.strip(): | |
| raise gr.Error("El prompt no puede estar vacio.") | |
| duracion = max(5, min(60, int(duracion))) | |
| print(f"Generando: {prompt} | objetivo {duracion}s") | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| modelo.to(device) | |
| # --- Pasada 1: generacion inicial solo con texto --- | |
| base_seg = min(duracion, 30) | |
| inputs = processor(text=[prompt], padding=True, return_tensors="pt").to(device) | |
| with torch.no_grad(): | |
| salida = modelo.generate(**inputs, max_new_tokens=min(int(base_seg * 50), MAX_TOKENS)) | |
| audio = _a_numpy(salida) | |
| # --- Pasadas de continuacion hasta llegar a la duracion --- | |
| while len(audio) < int(duracion * SR): | |
| restante_seg = duracion - len(audio) / SR | |
| nuevos_seg = min(CHUNK_MAX_SEG, restante_seg + 1) # +1 de margen para el cruce | |
| nuevos_tokens = int(nuevos_seg * 50) | |
| # cogemos la cola como arranque | |
| cola = audio[-int(CONTEXT_SEG * SR):] | |
| inputs = processor( | |
| audio=cola, | |
| sampling_rate=SR, | |
| text=[prompt], | |
| padding=True, | |
| return_tensors="pt", | |
| ).to(device) | |
| with torch.no_grad(): | |
| cont = modelo.generate(**inputs, max_new_tokens=nuevos_tokens) | |
| cont_audio = _a_numpy(cont) | |
| # la salida incluye la cola que le dimos: la quitamos y nos quedamos lo nuevo | |
| n_prompt = (len(cola) // FRAME) * FRAME | |
| parte_nueva = cont_audio[n_prompt:] | |
| if len(parte_nueva) < FRAME: # por si acaso no genero nada nuevo | |
| break | |
| audio = _crossfade_join(audio, parte_nueva, SR) | |
| audio = audio[: int(duracion * SR)] | |
| # --- Limpieza final: normalizar, seguridad, headroom, fade-out --- | |
| pico = np.max(np.abs(audio)) | |
| if pico > 0: | |
| audio = audio / pico | |
| audio = np.clip(audio, -1.0, 1.0) * 0.97 | |
| fade_len = min(len(audio), int(SR * 0.15)) | |
| if fade_len > 0: | |
| audio[-fade_len:] = audio[-fade_len:] * np.linspace(1.0, 0.0, fade_len) | |
| audio = (audio * 32767).astype(np.int16) | |
| tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) | |
| scipy.io.wavfile.write(tmp.name, SR, audio) | |
| print(f"Audio listo ({len(audio)/SR:.1f}s): {tmp.name}") | |
| return tmp.name | |
| with gr.Blocks(title="Fabrica de Musica") as demo: | |
| gr.Markdown("# Fabrica de Musica Teshua") | |
| gr.Markdown("Melodias continuas de hasta 60s con musicgen-medium (el modelo continua la composicion, no la repite).") | |
| with gr.Row(): | |
| with gr.Column(): | |
| prompt_input = gr.Textbox(label="Prompt en ingles", placeholder="mystical ambient music, spiritual, no vocals", lines=2) | |
| duracion_input = gr.Slider(label="Duracion segundos", minimum=5, maximum=60, value=30, step=1) | |
| btn = gr.Button("Generar Musica", variant="primary") | |
| with gr.Column(): | |
| audio_output = gr.Audio(label="Musica generada", type="filepath") | |
| btn.click(fn=generar_musica, inputs=[prompt_input, duracion_input], outputs=audio_output, api_name="generar_musica") | |
| demo.launch(server_name="0.0.0.0", server_port=7860) | |