ohfiftyb252 commited on
Commit
ef8bde4
·
verified ·
1 Parent(s): 9e06a06

Create trap_transformer.py

Browse files
Files changed (1) hide show
  1. trap_transformer.py +180 -0
trap_transformer.py ADDED
@@ -0,0 +1,180 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Trap Style Transformer - Convert MIDI to Trap Beat
4
+ Adds 808s, hi-hats, and trap patterns to existing MIDI
5
+ """
6
+
7
+ import pretty_midi as pm
8
+ import numpy as np
9
+ from pathlib import Path
10
+ import random
11
+
12
+ class TrapTransformer:
13
+ def __init__(self, bpm=140):
14
+ self.bpm = bpm
15
+ self.beat_duration = 60.0 / bpm
16
+
17
+ def extract_bass_notes(self, midi_data):
18
+ """Extract low notes (bass/808 candidate)"""
19
+ bass_notes = []
20
+ for track in midi_data.instruments:
21
+ for note in track.notes:
22
+ # Notes below C3 (~130Hz) are bass candidates
23
+ if note.pitch < 48: # C3
24
+ bass_notes.append(note)
25
+ return bass_notes
26
+
27
+ def create_808_bass_pattern(self, original_bass_notes, duration_seconds):
28
+ """Enhance bass with longer 808 sustain"""
29
+ enhanced_bass = []
30
+
31
+ for note in original_bass_notes:
32
+ # Make bass longer (808 characteristic)
33
+ new_note = pm.Note(
34
+ velocity=min(note.velocity + 20, 127), # Louder
35
+ pitch=note.pitch,
36
+ start=note.start,
37
+ end=min(note.end + 0.3, duration_seconds) # Extend by 300ms
38
+ )
39
+ enhanced_bass.append(new_note)
40
+
41
+ return enhanced_bass
42
+
43
+ def add_hi_hat_patterns(self, duration_seconds):
44
+ """Add trap hi-hat rolls (triplets/fast hats)"""
45
+ hat_track = pm.Instrument(program=0, is_drum=True)
46
+
47
+ # Hi-hat pattern: fast 1/32 notes with occasional rolls
48
+ hat_start = 0.0
49
+ while hat_start < duration_seconds:
50
+ # Regular 1/8 or 1/16 hi-hats
51
+ hat_duration = self.beat_duration / 4 # 1/16 notes
52
+
53
+ # Random chance for roll (1/32 or 1/64)
54
+ if random.random() < 0.3: # 30% chance of roll
55
+ roll_count = random.choice([4, 8, 16]) # Number of rolled hats
56
+ for i in range(roll_count):
57
+ hat = pm.Note(
58
+ velocity=random.randint(80, 120),
59
+ pitch=42, # Closed hi-hat
60
+ start=hat_start + (i * hat_duration / roll_count),
61
+ end=hat_start + ((i + 1) * hat_duration / roll_count)
62
+ )
63
+ hat_track.notes.append(hat)
64
+ hat_start += hat_duration
65
+ else:
66
+ hat = pm.Note(
67
+ velocity=random.randint(90, 115),
68
+ pitch=42, # Closed hi-hat
69
+ start=hat_start,
70
+ end=hat_start + hat_duration * 0.5
71
+ )
72
+ hat_track.notes.append(hat)
73
+ hat_start += hat_duration
74
+
75
+ return hat_track
76
+
77
+ def add_trap_drums(self, duration_seconds):
78
+ """Add 808 kick and snare/clap pattern"""
79
+ drum_track = pm.Instrument(program=0, is_drum=True)
80
+
81
+ # Kick pattern: every 1st and 3rd beat (typical trap half-time)
82
+ kick_times = [beat * self.beat_duration * 2 for beat in range(int(duration_seconds / (self.beat_duration * 2)))]
83
+
84
+ for kick_time in kick_times:
85
+ if kick_time < duration_seconds:
86
+ kick = pm.Note(
87
+ velocity=120, # Punchy kick
88
+ pitch=36, # Kick drum
89
+ start=kick_time,
90
+ end=kick_time + 0.1
91
+ )
92
+ drum_track.notes.append(kick)
93
+
94
+ # Snare/clap on beats 2 and 4
95
+ snare_times = [(beat + 1) * self.beat_duration * 2 for beat in range(int(duration_seconds / (self.beat_duration * 2)))]
96
+
97
+ for snare_time in snare_times:
98
+ if snare_time < duration_seconds:
99
+ snare = pm.Note(
100
+ velocity=110,
101
+ pitch=38, # Snare
102
+ start=snare_time,
103
+ end=snare_time + 0.08
104
+ )
105
+ drum_track.notes.append(snare)
106
+
107
+ return drum_track
108
+
109
+ def apply_sidechain_compression(self, midi_data, kick_times):
110
+ """Simulate sidechain pumping effect"""
111
+ # Lower melody volume slightly before each kick
112
+ for instrument in midi_data.instruments:
113
+ if not instrument.is_drum:
114
+ for note in instrument.notes:
115
+ # Find kicks that overlap with this note
116
+ for kick_time in kick_times:
117
+ if note.start <= kick_time <= note.end:
118
+ # Reduce velocity slightly
119
+ note.velocity = max(0, note.velocity - 10)
120
+ return midi_data
121
+
122
+ def transform_to_trap(self, input_midi_path, output_midi_path=None):
123
+ """Main function: Transform any MIDI to Trap style"""
124
+ midi_data = pm.PrettyMIDI(str(input_midi_path))
125
+
126
+ # Calculate duration
127
+ duration = midi_data.get_end_time()
128
+
129
+ # Extract and enhance bass
130
+ bass_notes = self.extract_bass_notes(midi_data)
131
+ enhanced_bass = self.create_808_bass_pattern(bass_notes, duration)
132
+
133
+ # Create new trap tracks
134
+ hat_track = self.add_hi_hat_patterns(duration)
135
+ drum_track = self.add_trap_drums(duration)
136
+
137
+ # Replace bass track with enhanced version
138
+ # (or add as new track depending on preference)
139
+
140
+ # Apply sidechain compression
141
+ kick_times = [n.start for n in drum_track.notes if n.pitch == 36]
142
+ midi_data = self.apply_sidechain_compression(midi_data, kick_times)
143
+
144
+ # Add trap tracks to the MIDI
145
+ midi_data.instruments.append(hat_track)
146
+ midi_data.instruments.append(drum_track)
147
+
148
+ # Update tempo to trap range
149
+ midi_data.tempos = [self.bpm]
150
+
151
+ # Save if output path provided
152
+ if output_midi_path:
153
+ midi_data.write(str(output_midi_path))
154
+
155
+ return midi_data
156
+
157
+ def transcribe_and_convert(audio_path, model_size="medium", trap_output_path=None):
158
+ """Full pipeline: Transcribe audio → Transform to Trap → Save MIDI"""
159
+ from muscriptor import TranscriptionModel
160
+
161
+ # Step 1: Transcribe with MuScriptor
162
+ print("🎹 Transcribing audio...")
163
+ model = TranscriptionModel.load_model(model_size)
164
+
165
+ # Get MIDI bytes from original transcription
166
+ midi_bytes = model.transcribe_to_midi(str(audio_path))
167
+
168
+ # Save temporary MIDI file
169
+ temp_midi = Path(tempfile.mktemp(suffix=".mid"))
170
+ temp_midi.write_bytes(midi_bytes)
171
+
172
+ # Step 2: Transform to Trap
173
+ print("🔥 Converting to Trap style...")
174
+ transformer = TrapTransformer(bpm=140)
175
+ trap_midi = transformer.transform_to_trap(temp_midi, trap_output_path)
176
+
177
+ # Cleanup
178
+ temp_midi.unlink()
179
+
180
+ return trap_midi