""" End-to-end tests for the music generation pipeline. Tests: tokenizer, model, generation, full pipeline. """ import sys from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) import torch import numpy as np from src.s01_config import ModelConfig, GenConfig from src.s02_tokenizer import MusicTokenizer, BOS_TOKEN, EOS_TOKEN, VOCAB_SIZE from src.s04_model import MusicTransformer from src.s06_generator import generate def test_tokenizer_roundtrip(): """Test tokenizer encode/decode produces valid events.""" tok = MusicTokenizer() assert tok.vocab_size == VOCAB_SIZE assert tok.bos_id == BOS_TOKEN assert tok.eos_id == EOS_TOKEN # Test individual token encoding note_on = tok.note_on_token(60) # Middle C event = tok.decode_token(note_on) assert event["type"] == "NoteOn" assert event["value"] == 60 vel = tok.velocity_token(100) event = tok.decode_token(vel) assert event["type"] == "Velocity" ts = tok.timeshift_token(500) # 500ms event = tok.decode_token(ts) assert event["type"] == "TimeShift" assert event["value"] == 500 print("PASS: test_tokenizer_roundtrip") def test_tokenizer_midi_conversion(): """Test MIDI to tokens and back.""" try: import pretty_midi except ImportError: print("SKIP: test_tokenizer_midi_conversion (pretty_midi not installed)") return tok = MusicTokenizer() # Create a simple test MIDI midi = pretty_midi.PrettyMIDI(initial_tempo=120) inst = pretty_midi.Instrument(program=0) # C major chord for pitch in [60, 64, 67]: note = pretty_midi.Note(velocity=80, pitch=pitch, start=0.0, end=1.0) inst.notes.append(note) # Second chord at 1.0s for pitch in [65, 69, 72]: note = pretty_midi.Note(velocity=90, pitch=pitch, start=1.0, end=2.0) inst.notes.append(note) midi.instruments.append(inst) # Tokenize tokens = tok.midi_to_tokens(midi) assert len(tokens) > 5 assert tokens[0] == BOS_TOKEN assert tokens[-1] == EOS_TOKEN # Decode back to MIDI midi_out = tok.tokens_to_midi(tokens) assert len(midi_out.instruments) == 1 assert len(midi_out.instruments[0].notes) > 0 print(f"PASS: test_tokenizer_midi_conversion ({len(tokens)} tokens, " f"{len(midi_out.instruments[0].notes)} notes)") def test_model_forward(): """Test model forward pass and loss computation.""" config = ModelConfig( vocab_size=VOCAB_SIZE, dim=64, n_layers=2, n_heads=4, n_kv_heads=2, max_seq_len=128, dropout=0.0, ) model = MusicTransformer.from_config(config) # Check parameter count is reasonable n_params = model.count_parameters() assert n_params > 0 print(f" Model params: {n_params:,}") # Forward pass batch_size = 2 seq_len = 32 input_ids = torch.randint(0, config.vocab_size, (batch_size, seq_len)) targets = torch.randint(0, config.vocab_size, (batch_size, seq_len)) logits, loss = model(input_ids, targets) assert logits.shape == (batch_size, seq_len, config.vocab_size) assert loss is not None assert loss.item() > 0 # Backward pass (check gradients flow) loss.backward() grad_norms = [p.grad.norm().item() for p in model.parameters() if p.grad is not None] assert len(grad_norms) > 0 assert all(not np.isnan(g) for g in grad_norms) print(f"PASS: test_model_forward (loss={loss.item():.4f})") def test_model_gradient_checkpoint(): """Test that gradient checkpointing works and reduces memory.""" config = ModelConfig( vocab_size=VOCAB_SIZE, dim=64, n_layers=4, n_heads=4, n_kv_heads=2, max_seq_len=128, dropout=0.0, ) model = MusicTransformer.from_config(config) model.grad_checkpoint = True input_ids = torch.randint(0, config.vocab_size, (2, 64)) targets = torch.randint(0, config.vocab_size, (2, 64)) logits, loss = model(input_ids, targets) loss.backward() assert loss.item() > 0 print(f"PASS: test_model_gradient_checkpoint (loss={loss.item():.4f})") def test_generation(): """Test autoregressive generation.""" config = ModelConfig( vocab_size=VOCAB_SIZE, dim=64, n_layers=2, n_heads=4, n_kv_heads=2, max_seq_len=128, dropout=0.0, ) model = MusicTransformer.from_config(config) tokenizer = MusicTokenizer() gen_config = GenConfig( temperature=0.8, top_k=20, top_p=0.9, max_tokens=50, repetition_penalty=1.1, seed=42, ) tokens = generate(model, tokenizer, gen_config, device=torch.device("cpu")) assert len(tokens) > 1 assert tokens[0] == BOS_TOKEN print(f"PASS: test_generation ({len(tokens)} tokens generated)") def test_generation_to_midi(): """Test full pipeline: generate tokens → convert to MIDI.""" try: import pretty_midi except ImportError: print("SKIP: test_generation_to_midi (pretty_midi not installed)") return config = ModelConfig( vocab_size=VOCAB_SIZE, dim=64, n_layers=2, n_heads=4, n_kv_heads=2, max_seq_len=128, dropout=0.0, ) model = MusicTransformer.from_config(config) tokenizer = MusicTokenizer() gen_config = GenConfig( temperature=1.0, top_k=50, top_p=0.95, max_tokens=100, repetition_penalty=1.1, seed=123, ) tokens = generate(model, tokenizer, gen_config, device=torch.device("cpu")) midi = tokenizer.tokens_to_midi(tokens) assert midi is not None assert len(midi.instruments) == 1 print(f"PASS: test_generation_to_midi ({len(tokens)} tokens → " f"{len(midi.instruments[0].notes)} notes)") def test_dataset_creation(): """Test MidiTokenDataset with synthetic data.""" from src.s03_dataset import MidiTokenDataset # Synthetic token sequences sequences = [ [BOS_TOKEN] + list(np.random.randint(4, VOCAB_SIZE, size=100)) + [EOS_TOKEN] for _ in range(20) ] ds = MidiTokenDataset(sequences, max_seq_len=64, pad_id=0) assert len(ds) == 20 input_ids, targets = ds[0] assert input_ids.shape == (64,) assert targets.shape == (64,) assert input_ids.dtype == torch.long print(f"PASS: test_dataset_creation ({len(ds)} sequences)") if __name__ == "__main__": print("=" * 60) print("MUSIC GENERATION LLM — TESTS") print("=" * 60) tests = [ test_tokenizer_roundtrip, test_tokenizer_midi_conversion, test_model_forward, test_model_gradient_checkpoint, test_generation, test_generation_to_midi, test_dataset_creation, ] passed = 0 failed = 0 skipped = 0 for test in tests: try: test() passed += 1 except Exception as e: if "SKIP" in str(e): skipped += 1 else: print(f"FAIL: {test.__name__}: {e}") import traceback traceback.print_exc() failed += 1 print("=" * 60) print(f"Results: {passed} passed, {failed} failed, {skipped} skipped") print("=" * 60) sys.exit(1 if failed > 0 else 0)