Muscript / ai_style_transfer.py
ohfiftyb252's picture
Create ai_style_transfer.py
612ec05 verified
Raw History Blame Contribute Delete
8.19 kB
#!/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