#!/usr/bin/env python3 """ AI Style Transfer Module - Uses MusicGen/Riffusion to reimagine MIDI """ import subprocess from pathlib import Path from typing import Optional import tempfile class AIStyleTransfer: def __init__(self, provider="huggingface-api"): """ provider options: - "local" : Run MusicGen locally (GPU required) - "huggingface-api" : Use HF Inference API (free tier available) """ self.provider = provider def extract_tempo(self, audio_path): """Estimate tempo from audio using librosa""" try: import librosa y, sr = librosa.load(str(audio_path), duration=10) # Load first 10s tempo, _ = librosa.beat.beat_track(y=y, sr=sr) return float(tempo) except ImportError: return 140.0 # Default trap BPM def create_trap_prompt(self, instruments=None, mood="dark", intensity="aggressive"): """Build optimized prompt for MusicGen/Riffusion""" base_elements = [ f"{mood} trap beat", "heavy 808 bass", "fast hi-hat rolls", f"{intensity} energy", "140 bpm" ] if instruments: # Add melodic elements based on original transcription base_elements.append("melodic synth lead") prompt = ", ".join(base_elements) return prompt def musicgen_local(self, audio_description: str, duration: int = 30): """ Run MusicGen locally via transformers Requires GPU with 16GB+ VRAM """ try: from transformers import AutoProcessor, AutoModelForTextToWaveform import torch processor = AutoProcessor.from_pretrained("facebook/musicgen-large") model = AutoModelForTextToWaveform.from_pretrained( "facebook/musicgen-large", torch_dtype=torch.float16 ).to("cuda") inputs = processor( text=audio_description, padding=True, return_tensors="pt" ) with torch.no_grad(): audio_values = model.generate( **inputs.to("cuda"), max_new_tokens=256, guidance_scale=3, do_sample=True, temperature=1.0 ) # Convert tensor to WAV file sampling_rate = model.config.sampling_rate audio_array = audio_values[0, 0].cpu().numpy() from scipy.io.wavfile import write temp_wav = Path(tempfile.mktemp(suffix=".wav")) write(temp_wav, sampling_rate, (audio_array * 32767).astype(np.int16)) return temp_wav except Exception as e: print(f"MusicGen local failed: {e}") return None def riffusion_hf_api(self, prompt: str, seed_image: str = None): """ Use Riffusion via Hugging Face Inference API No GPU needed, works on free tier """ try: from huggingface_hub import InferenceClient client = InferenceClient( token=os.environ.get("HF_TOKEN", ""), model="riffusion/riffusion" ) # Generate spectrogram-based audio output = client.text_to_sound( prompt=prompt, duration=30 # seconds ) # Save as WAV temp_wav = Path(tempfile.mktemp(suffix=".wav")) temp_wav.write_bytes(output) return temp_wav except Exception as e: print(f"Riffusion failed: {e}") return None def musicgen_hf_api(self, prompt: str, duration: int = 30): """ Use MusicGen via Hugging Face Inference API Slower but no local GPU required """ try: from huggingface_hub import InferenceClient import time client = InferenceClient(token=os.environ.get("HF_TOKEN", "")) # Start generation task task = client.text_to_audio( inputs=prompt, parameters={ "model": "facebook/musicgen-large", "duration": duration } ) # Poll for completion max_retries = 60 for attempt in range(max_retries): status = client.check_task_status(task.task_id) if status["status"] == "completed": return Path(status["output"]["audio_path"]) elif status["status"] == "failed": print(f"Generation failed: {status}") break time.sleep(2) return None except Exception as e: print(f"MusicGen API failed: {e}") return None def stem_separation(self, audio_path): """ Separate vocals/instruments using Demucs Useful for isolating melody before style transfer """ try: from demucs.apply import apply_model from demucs.pretrained import get_model import torchaudio # Load pre-trained Demucs model model = get_model("htdemucs") # Load audio wav, sr = torchaudio.load(str(audio_path)) # Apply separation stems = apply_model(model, wav) # Save separated tracks (vocals, drums, bass, other) stem_names = model.stems output_dir = Path(tempfile.mkdtemp()) for i, name in enumerate(stem_names): stem_path = output_dir / f"{name}.wav" torchaudio.save(str(stem_path), stems[:, :, i:i+1], sr) return output_dir except Exception as e: print(f"Stem separation failed: {e}") return None def mix_trap_final(self, original_midi_path, ai_generated_wav, output_path): """ Mix original MIDI with AI-generated Trap backing track Creates final composite track """ try: from pydub import AudioSegment import pretty_midi as pm # Load AI-generated trap beat trap_beat = AudioSegment.from_wav(str(ai_generated_wav)) # Render MIDI to audio (using FluidSynth or similar) midi_audio = self.render_midi_to_audio(original_midi_path) # Align and mix mixed = trap_beat.overlay(midi_audio) # Export final mixed.export(str(output_path), format="mp3", bitrate="320k") return output_path except Exception as e: print(f"Mixing failed: {e}") return None def render_midi_to_audio(self, midi_path, soundfont="/tmp/default.sf2"): """Render MIDI file to audio using FluidSynth""" try: import fluidsynth import io fs = fluidsynth.Synth() fs.start(driver="disk") sfid = fs.sfload(soundfont) fs.program_select(0, sfid, 0, 0) midi_data = pm.PrettyMIDI(str(midi_path)) midi_data.fluidsynth(fs, sfid) # Get rendered audio audio = fs.get_samples() # Save to file from scipy.io.wavfile import write temp_wav = Path(tempfile.mktemp(suffix=".wav")) write(temp_wav, 48000, audio.T.astype('int16')) return temp_wav except Exception as e: print(f"MIDI rendering failed: {e}") return None