""" MIDI Batch Renderer using DawDreamer with Multiprocessing Renders MIDI files to audio using VST instruments with random presets Based on DawDreamer's parallel_render.py example """ import dawdreamer as daw from scipy.io import wavfile import os import glob import argparse import json from pathlib import Path import mido import time import logging import multiprocessing import traceback from collections import namedtuple from tqdm import tqdm import random import subprocess import tempfile # Setup logging logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) # Named tuple for work items RenderItem = namedtuple("RenderItem", "midi_file output_file preset_path instrument_type") class MidiInfo: """Extract basic info from MIDI files""" def __init__(self, midi_file_path): self.path = Path(midi_file_path) self._load_midi() def _load_midi(self): """Load MIDI file and extract tempo/duration""" try: midi_file = mido.MidiFile(str(self.path)) self.length = midi_file.length self.bpm = self._extract_bpm(midi_file) self.midi_file_path = Path(str(self.path)) except Exception as e: logger.warning(f"Failed to load MIDI info for {self.path}: {e}") self.length = 30.0 # Default fallback self.bpm = 120.0 # Default fallback def _extract_bpm(self, midi_file): """Extract BPM from MIDI file""" for track in midi_file.tracks: for msg in track: if msg.type == "set_tempo": return mido.tempo2bpm(msg.tempo) return 120.0 class PresetManager: """Manages VST presets organized by instrument type""" def __init__(self, preset_root_dir): self.preset_root_dir = Path(preset_root_dir) self.instrument_presets = {} self._scan_presets() def _scan_presets(self): """Scan preset directory and organize by instrument type""" if not self.preset_root_dir.exists(): raise FileNotFoundError(f"Preset directory not found: {self.preset_root_dir}") logger.info(f"Scanning presets in: {self.preset_root_dir}") for instrument_dir in self.preset_root_dir.iterdir(): if instrument_dir.is_dir(): instrument_name = instrument_dir.name.lower() preset_files = list(instrument_dir.glob("*.fxp")) + list(instrument_dir.glob("*.fxb")) if preset_files: self.instrument_presets[instrument_name] = preset_files logger.info(f" Found {len(preset_files)} presets for '{instrument_name}'") if not self.instrument_presets: raise ValueError(f"No presets found in {self.preset_root_dir}") logger.info(f"Total instrument types: {len(self.instrument_presets)}") def get_random_preset(self, instrument_type): """Get a random preset for the specified instrument type""" instrument_type = instrument_type.lower() # if instrument_type == "other": # # For "other", choose a random instrument category # chosen_instrument = random.choice(list(self.instrument_presets.keys())) # return random.choice(self.instrument_presets[chosen_instrument]) if instrument_type in self.instrument_presets: return random.choice(self.instrument_presets[instrument_type]) raise ValueError(f"No presets found for instrument type: {instrument_type}") def should_render_file(self, midi_filename): """Check if MIDI file should be rendered based on suffix""" stem = Path(midi_filename).stem.lower() # Check for _other suffix # if stem.endswith('_other'): # return True, 'other' # Check for instrument-specific suffixes for instrument_type in self.instrument_presets.keys(): if stem.endswith(f"_{instrument_type}"): return True, instrument_type return False, None def get_available_instruments(self): """Get list of available instrument types""" return list(self.instrument_presets.keys()) class Worker: """Worker process for rendering MIDI files""" def __init__( self, queue: multiprocessing.Queue, vst_path: str, sample_rate=48000, buffer_size=512, output_dir="output", opus_bitrate=128, opus_vbr=True, ): self.queue = queue self.vst_path = vst_path self.sample_rate = sample_rate self.buffer_size = buffer_size self.output_dir = Path(output_dir) self.opus_bitrate = opus_bitrate self.opus_vbr = opus_vbr # Results will be stored here self.results = [] def startup(self): """Initialize DawDreamer engine for this worker""" try: self.engine = daw.RenderEngine(self.sample_rate, self.buffer_size) self.synth = self.engine.make_plugin_processor("synth", self.vst_path) self.engine.load_graph([(self.synth, [])]) logger.debug(f"Worker initialized with VST: {self.vst_path}") except Exception as e: logger.error(f"Failed to initialize worker: {e}") raise def save_as_opus_ffmpeg(self, audio, output_path, sample_rate, bitrate=128, vbr=True): """Save audio as Opus using FFmpeg""" # Create temporary WAV file with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_wav: tmp_wav_path = tmp_wav.name # Save as WAV first wavfile.write(tmp_wav_path, sample_rate, audio.transpose()) # Build FFmpeg command cmd = [ "ffmpeg", "-i", tmp_wav_path, "-c:a", "libopus", "-b:a", f"{bitrate}k", "-y", # Overwrite output files str(output_path), ] # Add VBR settings if enabled if vbr: cmd.extend(["-vbr", "on"]) try: # Run FFmpeg result = subprocess.run(cmd, check=True, capture_output=True, text=True) except subprocess.CalledProcessError as e: logger.error(f"FFmpeg error: {e.stderr}") raise finally: # Clean up temporary WAV file try: os.unlink(tmp_wav_path) except OSError: pass def process_item(self, item: RenderItem): """Process a single MIDI file rendering""" try: # Load preset self.synth.load_preset(str(item.preset_path)) # Get preset name for output filename preset_name = Path(item.preset_path).stem output_path = Path(item.output_file) final_output = output_path.parent / f"{output_path.stem}.opus" # Load MIDI info midi_data = MidiInfo(item.midi_file) duration = midi_data.length + 1.0 # Add tail self.engine.set_bpm(midi_data.bpm) # Load MIDI file the_path = str(item.midi_file) self.synth.load_midi(the_path, clear_previous=True, beats=False) # Render start_time = time.time() self.engine.render(duration) self.synth.clear_midi() render_time = time.time() - start_time # Get audio audio = self.engine.get_audio() # Create output directory os.makedirs(final_output.parent, exist_ok=True) # Save as Opus using FFmpeg self.save_as_opus_ffmpeg( audio, final_output, self.sample_rate, self.opus_bitrate, self.opus_vbr ) # Record result result = { "midi_file": str(item.midi_file), "output_file": str(final_output), "duration": duration, "bpm": midi_data.bpm, "render_time": render_time, "real_time_ratio": duration / render_time if render_time > 0 else 0, "instrument_type": item.instrument_type, "preset_used": str(item.preset_path), "opus_bitrate": self.opus_bitrate, "success": True, } self.results.append(result) except Exception as e: logger.error(f"Failed to process {item.midi_file}: {e}") result = { "midi_file": str(item.midi_file), "output_file": str(item.output_file), "error": str(e), "success": False, } self.results.append(result) def run(self): """Main worker loop - processes items from queue until empty""" try: self.startup() while True: try: item = self.queue.get_nowait() self.process_item(item) except multiprocessing.queues.Empty: break return self.results except Exception as e: logger.error(f"Worker exception: {e}") return traceback.format_exc() class MIDIBatchRenderer: """Batch MIDI renderer using DawDreamer's multiprocessing pattern""" def __init__(self, vst_path, sample_rate=48000, buffer_size=512, opus_bitrate=128, opus_vbr=True): self.vst_path = vst_path self.sample_rate = sample_rate self.buffer_size = buffer_size self.opus_bitrate = opus_bitrate self.opus_vbr = opus_vbr if not os.path.exists(vst_path): raise FileNotFoundError(f"VST not found: {vst_path}") def find_midi_files(self, midi_dir): """Find all MIDI files recursively""" midi_files = [] for pattern in ["*.mid", "*.midi", "**/*.mid", "**/*.midi"]: midi_files.extend(glob.glob(os.path.join(midi_dir, pattern), recursive=True)) return list(set(midi_files)) # Remove duplicates def prepare_render_items(self, midi_files, midi_dir, output_dir, preset_manager): """Prepare work items for the queue""" render_items = [] skipped_files = [] for midi_file in midi_files: # Check if file should be rendered should_render, instrument_type = preset_manager.should_render_file(midi_file) if not should_render: skipped_files.append(midi_file) continue # Get random preset try: preset_path = preset_manager.get_random_preset(instrument_type) except ValueError as e: logger.warning(f"Skipping {midi_file}: {e}") skipped_files.append(midi_file) continue # Prepare output path (preserving directory structure) rel_path = os.path.relpath(midi_file, midi_dir) output_file = os.path.join(output_dir, Path(rel_path).with_suffix(".opus")) # Create render item item = RenderItem( midi_file=midi_file, output_file=output_file, preset_path=preset_path, instrument_type=instrument_type, ) render_items.append(item) return render_items, skipped_files def batch_render_parallel(self, midi_dir, output_dir, preset_manager, num_workers=None): """Batch render MIDI files using multiprocessing""" logger.info(f"Starting parallel batch render (Opus output)...") logger.info(f"MIDI dir: {midi_dir} | Output dir: {output_dir}") logger.info(f"Opus settings: {self.opus_bitrate}kbps, VBR: {self.opus_vbr}") logger.info(f"Available instruments: {preset_manager.get_available_instruments()}") # Find MIDI files midi_files = self.find_midi_files(midi_dir) if not midi_files: logger.warning(f"No MIDI files found in: {midi_dir}") return [] # Prepare work items render_items, skipped_files = self.prepare_render_items( midi_files, midi_dir, output_dir, preset_manager ) logger.info( f"Found {len(midi_files)} files | Renderable: {len(render_items)} | Skipped: {len(skipped_files)}" ) if not render_items: logger.warning("No files match instrument naming patterns!") return [] # Create output directory os.makedirs(output_dir, exist_ok=True) # Determine number of workers num_processes = num_workers or multiprocessing.cpu_count() logger.info(f"Using {num_processes} worker processes") # Create queue and add items input_queue = multiprocessing.Manager().Queue() for item in render_items: input_queue.put(item) # Start processing start_time = time.time() all_results = [] # Create multiprocessing Pool following DawDreamer's pattern with multiprocessing.Pool(processes=num_processes) as pool: # Create workers workers = [] for i in range(num_processes): worker = Worker( input_queue, self.vst_path, self.sample_rate, self.buffer_size, output_dir, self.opus_bitrate, self.opus_vbr, ) async_result = pool.apply_async(worker.run) workers.append(async_result) # Progress tracking pbar = tqdm(total=len(render_items), desc="Rendering MIDI files to Opus") completed_last = 0 while True: incomplete_count = sum(1 for w in workers if not w.ready()) if incomplete_count == 0: break # Update progress bar # Estimate progress based on queue size (approximate) try: remaining = input_queue.qsize() completed = len(render_items) - remaining pbar.update(completed - completed_last) completed_last = completed except: pass time.sleep(0.1) pbar.close() # Collect results from all workers for i, worker_result in enumerate(workers): try: worker_results = worker_result.get() if isinstance(worker_results, str): # Error traceback logger.error(f"Worker {i} exception:\n{worker_results}") else: all_results.extend(worker_results) except Exception as e: logger.error(f"Failed to get results from worker {i}: {e}") # Log results total_time = time.time() - start_time successful = [r for r in all_results if r.get("success")] failed = [r for r in all_results if not r.get("success")] self._log_results(all_results, successful, failed, total_time) return all_results def _log_results(self, results, successful, failed, total_time): """Log processing results""" logger.info(f"\n=== Batch Complete ===") logger.info(f"Total: {len(results)} | Success: {len(successful)} | Failed: {len(failed)}") logger.info(f"Total time: {total_time:.2f}s") if successful: avg_ratio = sum(r["real_time_ratio"] for r in successful) / len(successful) total_duration = sum(r["duration"] for r in successful) logger.info(f"Avg real-time ratio: {avg_ratio:.1f}x | Audio rendered: {total_duration:.1f}s") # Instrument distribution inst_counts = {} for r in successful: inst_type = r.get("instrument_type", "unknown") inst_counts[inst_type] = inst_counts.get(inst_type, 0) + 1 logger.info(f"By instrument: {inst_counts}") if failed: logger.warning(f"Failed files: {[os.path.basename(r['midi_file']) for r in failed[:5]]}") def save_report(self, results, report_path): """Save processing report as JSON""" os.makedirs(os.path.dirname(report_path), exist_ok=True) with open(report_path, "w") as f: json.dump( { "summary": { "total_files": len(results), "successful": len([r for r in results if r.get("success")]), "failed": len([r for r in results if not r.get("success")]), "timestamp": time.strftime("%Y-%m-%d %H:%M:%S"), }, "results": results, }, f, indent=2, ) logger.info(f"Report saved: {report_path}") def main(): parser = argparse.ArgumentParser( description="Batch render MIDI files with VST instruments and random presets" ) parser.add_argument("vst_path", help="Path to VST instrument") parser.add_argument("midi_dir", help="Directory containing MIDI files") parser.add_argument("output_dir", help="Output directory for audio files") parser.add_argument("preset_dir", help="Directory with preset folders by instrument type") parser.add_argument("--sample-rate", type=int, default=48000, help="Sample rate (default: 48000)") parser.add_argument("--buffer-size", type=int, default=512, help="Buffer size (default: 512)") parser.add_argument( "--num-workers", type=int, default=None, help="Number of workers (default: CPU count)" ) parser.add_argument("--report", help="Save report to JSON file") parser.add_argument("--verbose", "-v", action="store_true", help="Verbose logging") args = parser.parse_args() if args.verbose: logging.getLogger().setLevel(logging.DEBUG) # Validate inputs for path, name in [ (args.vst_path, "VST"), (args.midi_dir, "MIDI directory"), (args.preset_dir, "Preset directory"), ]: if not os.path.exists(path): logger.error(f"{name} not found: {path}") return 1 try: preset_manager = PresetManager(args.preset_dir) renderer = MIDIBatchRenderer(args.vst_path, args.sample_rate, args.buffer_size) results = renderer.batch_render_parallel( args.midi_dir, args.output_dir, preset_manager, args.num_workers ) if args.report: renderer.save_report(results, args.report) return 0 except Exception as e: logger.error(f"Error: {e}") return 1 if __name__ == "__main__": exit(main())