Muscript / trap_transformer.py
ohfiftyb252's picture
Update trap_transformer.py
495ab46 verified
Raw History Blame Contribute Delete
4.45 kB
#!/usr/bin/env python3
"""Trap Style MIDI Transformer"""
import pretty_midi as pm
import numpy as np
import random
from pathlib import Path
class TrapTransformer:
def __init__(self, bpm=140):
self.bpm = bpm
self.beat_duration = 60.0 / bpm
def extract_bass_notes(self, midi_data):
"""Pull low notes that can become 808s"""
bass_notes = []
for track in midi_data.instruments:
if track.is_drum:
continue
for note in track.notes:
if note.pitch < 48: # Below C3
bass_notes.append(note)
return bass_notes
def create_808_bass(self, bass_notes, duration_seconds):
"""Longer sustain + louder = 808 feel"""
enhanced = []
for note in bass_notes:
new_note = pm.Note(
velocity=min(note.velocity + 20, 127),
pitch=note.pitch,
start=note.start,
end=min(note.end + 0.3, duration_seconds),
)
enhanced.append(new_note)
return enhanced
def add_hi_hats(self, duration_seconds):
"""Fast hi-hat rolls (trap signature)"""
hat_track = pm.Instrument(program=0, is_drum=True, name="HiHats")
t = 0.0
while t < duration_seconds:
step = self.beat_duration / 4 # 1/16 notes
if random.random() < 0.3: # 30% chance β†’ roll
rolls = random.choice([4, 8])
for i in range(rolls):
hat_track.notes.append(
pm.Note(
velocity=random.randint(80, 120),
pitch=42,
start=t + (i * step / rolls),
end=t + ((i + 1) * step / rolls),
)
)
else:
hat_track.notes.append(
pm.Note(
velocity=random.randint(90, 115),
pitch=42,
start=t,
end=t + step * 0.5,
)
)
t += step
return hat_track
def add_trap_drums(self, duration_seconds):
"""Half-time kicks + snares"""
drum_track = pm.Instrument(program=0, is_drum=True, name="TrapDrums")
bar = self.beat_duration * 2 # Half-time bar
num_bars = int(duration_seconds / bar)
for b in range(num_bars):
kick_t = b * bar
snare_t = b * bar + self.beat_duration # Snare on beat 3 (half-time)
drum_track.notes.append(
pm.Note(velocity=120, pitch=36, start=kick_t, end=kick_t + 0.1)
)
drum_track.notes.append(
pm.Note(velocity=110, pitch=38, start=snare_t, end=snare_t + 0.08)
)
return drum_track
def apply_sidechain(self, midi_data, kick_times):
"""Simulate sidechain pump β€” duck melody before kicks"""
for instrument in midi_data.instruments:
if instrument.is_drum:
continue
for note in instrument.notes:
for kt in kick_times:
if note.start <= kt <= note.end:
note.velocity = max(1, note.velocity - 10)
return midi_data
def transform_to_trap(self, input_midi_path, output_midi_path=None):
"""Main: transform any MIDI β†’ Trap"""
midi_data = pm.PrettyMIDI(str(input_midi_path))
duration = midi_data.get_end_time()
# Enhance existing bass
bass_notes = self.extract_bass_notes(midi_data)
enhanced_bass = self.create_808_bass(bass_notes, duration)
# Add new drum + hat tracks
hat_track = self.add_hi_hats(duration)
drum_track = self.add_trap_drums(duration)
# Sidechain
kick_times = [n.start for n in drum_track.notes if n.pitch == 36]
midi_data = self.apply_sidechain(midi_data, kick_times)
# Append trap layers
midi_data.instruments.append(hat_track)
midi_data.instruments.append(drum_track)
# Set tempo
try:
for tempo_change in midi_data._tempo_changes:
tempo_change.tempo = self.bpm
except Exception:
pass
if output_midi_path:
midi_data.write(str(output_midi_path))
return midi_data