Spaces:
Running on Zero
Running on Zero
File size: 8,192 Bytes
612ec05 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 | #!/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 |