| """ |
| Data pipeline: downloads MIDI dataset, tokenizes, creates PyTorch DataLoaders. |
| Uses HuggingFace datasets for efficient streaming + caching. |
| Memory-efficient: processes files lazily, doesn't hold entire dataset in RAM. |
| """ |
| import logging |
| import os |
| import pickle |
| import signal |
| from pathlib import Path |
| from typing import Optional |
|
|
| import numpy as np |
| import torch |
| from torch.utils.data import Dataset, DataLoader |
|
|
|
|
| class _TimeoutError(Exception): |
| pass |
|
|
|
|
| def _timeout_handler(signum, frame): |
| raise _TimeoutError("File processing timed out") |
|
|
| from src.s01_config import DataConfig, PathConfig, TrainConfig |
| from src.s02_tokenizer import MusicTokenizer |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| class MidiTokenDataset(Dataset): |
| """ |
| PyTorch dataset of pre-tokenized MIDI sequences. |
| Stores token IDs as memory-mapped numpy arrays for RAM efficiency. |
| """ |
|
|
| def __init__(self, token_sequences: list[list[int]], max_seq_len: int, pad_id: int = 0): |
| self.max_seq_len = max_seq_len |
| self.pad_id = pad_id |
| |
| self.sequences = [s for s in token_sequences if len(s) >= 10] |
| logger.info(f"Dataset: {len(self.sequences)} sequences, max_len={max_seq_len}") |
|
|
| def __len__(self): |
| return len(self.sequences) |
|
|
| def __getitem__(self, idx): |
| seq = self.sequences[idx] |
|
|
| |
| if len(seq) > self.max_seq_len + 1: |
| |
| start = np.random.randint(0, len(seq) - self.max_seq_len) |
| seq = seq[start : start + self.max_seq_len + 1] |
|
|
| input_ids = seq[:-1] |
| target_ids = seq[1:] |
|
|
| |
| pad_len = self.max_seq_len - len(input_ids) |
| if pad_len > 0: |
| input_ids = input_ids + [self.pad_id] * pad_len |
| target_ids = target_ids + [self.pad_id] * pad_len |
|
|
| return ( |
| torch.tensor(input_ids, dtype=torch.long), |
| torch.tensor(target_ids, dtype=torch.long), |
| ) |
|
|
|
|
| def download_and_tokenize( |
| data_config: DataConfig, |
| path_config: PathConfig, |
| tokenizer: MusicTokenizer, |
| ) -> tuple[list[list[int]], list[list[int]]]: |
| """ |
| Download MIDI dataset from HuggingFace and tokenize all files. |
| Returns (train_sequences, val_sequences). |
| Caches tokenized data to disk for fast reload. |
| """ |
| cache_path = path_config.data_dir / "tokenized_cache.pkl" |
|
|
| if cache_path.exists(): |
| logger.info("Loading tokenized data from cache...") |
| with open(cache_path, "rb") as f: |
| data = pickle.load(f) |
| return data["train"], data["val"] |
|
|
| logger.info(f"Downloading dataset: {data_config.dataset_name}") |
| import pretty_midi |
| import subprocess |
| import glob |
|
|
| |
| midi_dir = path_config.data_dir / "midi_files" |
| if not midi_dir.exists(): |
| logger.info("Cloning MIDI dataset repo (faster than per-file download)...") |
| subprocess.run( |
| ["git", "clone", "--depth", "1", |
| f"https://huggingface.co/datasets/{data_config.dataset_name}", |
| str(midi_dir)], |
| check=True, |
| ) |
| else: |
| logger.info(f"Using cached MIDI files from {midi_dir}") |
|
|
| |
| midi_files = sorted( |
| glob.glob(f"{midi_dir}/**/*.mid", recursive=True) |
| + glob.glob(f"{midi_dir}/**/*.midi", recursive=True) |
| + glob.glob(f"{midi_dir}/**/*.MID", recursive=True) |
| ) |
| logger.info(f"Found {len(midi_files)} MIDI files") |
|
|
| all_sequences = [] |
| errors = 0 |
|
|
| for i, midi_path in enumerate(midi_files): |
| try: |
| |
| if os.path.getsize(midi_path) > 100_000: |
| errors += 1 |
| continue |
| |
| with open(midi_path, "rb") as f: |
| header = f.read(20) |
| if header.startswith(b"version https://git"): |
| errors += 1 |
| continue |
| |
| old_handler = signal.signal(signal.SIGALRM, _timeout_handler) |
| signal.alarm(30) |
| try: |
| midi = pretty_midi.PrettyMIDI(midi_path) |
| tokens = tokenizer.midi_to_tokens(midi, max_len=data_config.max_seq_len + 1) |
| if len(tokens) >= 20: |
| all_sequences.append(tokens) |
| finally: |
| signal.alarm(0) |
| signal.signal(signal.SIGALRM, old_handler) |
| except (_TimeoutError, Exception) as e: |
| errors += 1 |
| if errors <= 10: |
| logger.warning(f"Error processing {midi_path}: {e}") |
|
|
| if (i + 1) % 200 == 0: |
| logger.info(f" Processed {i+1}/{len(midi_files)}, valid={len(all_sequences)}, errors={errors}") |
|
|
| logger.info(f"Tokenization complete: {len(all_sequences)} sequences, {errors} errors") |
|
|
| |
| np.random.seed(42) |
| indices = np.random.permutation(len(all_sequences)) |
| split = int(len(all_sequences) * data_config.train_split) |
|
|
| train_seqs = [all_sequences[i] for i in indices[:split]] |
| val_seqs = [all_sequences[i] for i in indices[split:]] |
|
|
| |
| with open(cache_path, "wb") as f: |
| pickle.dump({"train": train_seqs, "val": val_seqs}, f) |
| logger.info(f"Cached tokenized data: train={len(train_seqs)}, val={len(val_seqs)}") |
|
|
| return train_seqs, val_seqs |
|
|
|
|
| def create_dataloaders( |
| data_config: DataConfig, |
| train_config: TrainConfig, |
| path_config: PathConfig, |
| tokenizer: MusicTokenizer, |
| ) -> tuple[DataLoader, DataLoader]: |
| """Create train and validation DataLoaders.""" |
| train_seqs, val_seqs = download_and_tokenize(data_config, path_config, tokenizer) |
|
|
| train_ds = MidiTokenDataset(train_seqs, data_config.max_seq_len, tokenizer.pad_id) |
| val_ds = MidiTokenDataset(val_seqs, data_config.max_seq_len, tokenizer.pad_id) |
|
|
| train_loader = DataLoader( |
| train_ds, |
| batch_size=train_config.batch_size, |
| shuffle=True, |
| num_workers=train_config.num_workers, |
| pin_memory=train_config.pin_memory, |
| prefetch_factor=train_config.prefetch_factor, |
| drop_last=True, |
| ) |
| val_loader = DataLoader( |
| val_ds, |
| batch_size=train_config.batch_size, |
| shuffle=False, |
| num_workers=train_config.num_workers, |
| pin_memory=train_config.pin_memory, |
| prefetch_factor=train_config.prefetch_factor, |
| drop_last=False, |
| ) |
|
|
| return train_loader, val_loader |
|
|