Spaces:
Running on Zero
Running on Zero
Download ai_style_transfer.py from ohfiftyb252/Muscript: direct link, hf CLI and curl.
- Browser
- Download file 8.19 kB
-
https://huggingface.co/spaces/ohfiftyb252/Muscript/resolve/main/ai_style_transfer.py
- Command line
-
hf download hf://spaces/ohfiftyb252/Muscript/ai_style_transfer.py
-
curl -L -o ai_style_transfer.py https://huggingface.co/spaces/ohfiftyb252/Muscript/resolve/main/ai_style_transfer.py
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 |