import unittest import numpy as np import random import sys import os from unittest.mock import Mock, patch, MagicMock # Add parent directory to path for imports sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from data_utils import song_to_samples, get_samples_for_song from data_types import SamplingParams, SampleData from block_types import SampleBlockType from modules.gpt import GPTConfig class TestToMusic(unittest.TestCase): def setUp(self): """Set up test fixtures with realistic data structures.""" random.seed(42) np.random.seed(42) # Mock configuration self.cfg = Mock() self.cfg.semantic_rate_hz = 25 self.cfg.semantic_codebook_size = 4096 # Mock song data and metadata self.song_data = np.random.randint(0, 4096, (1, 1000)) # 40 seconds at 25Hz self.song_meta = { "duration": 40.0, "sample_idx": 123, "artist": "Test Artist", "tags": "rock, guitar", } # Mock sampling parameters self.sampling_params = SamplingParams(allow_sample=True, prob_sample=0.1) # Mock sample data with required parameters self.sample_data = SampleData( data_row=self.song_data, data_meta=self.song_meta, sampling_params=self.sampling_params, audio_sample_tracks=[np.random.randint(0, 4096, (1, 250))], # 10 seconds at 25Hz ) def test_song_to_samples_basic(self): """Test basic song_to_samples functionality.""" samples = song_to_samples( self.song_data, self.song_meta, self.cfg, n_samples=2, sample_duration_range=(5, 15) ) self.assertEqual(len(samples), 2) for sample in samples: self.assertIsInstance(sample, dict) self.assertIn("data", sample) self.assertIn("meta", sample) self.assertIsInstance(sample["data"], np.ndarray) self.assertEqual(sample["data"].shape[0], 1) # Single codebook # Check sample duration is within range (5-15 seconds at 25Hz) duration_tokens = sample["data"].shape[1] duration_seconds = duration_tokens / self.cfg.semantic_rate_hz self.assertGreaterEqual(duration_seconds, 5) self.assertLessEqual(duration_seconds, 15) def test_song_to_samples_edge_cases(self): """Test song_to_samples with edge cases.""" # Test with very short song short_song = np.random.randint(0, 4096, (1, 100)) # 4 seconds short_meta = {"duration": 4.0, "sample_idx": 456} samples = song_to_samples(short_song, short_meta, self.cfg, n_samples=1) # Short songs might return no samples if they're too short for the minimum duration self.assertGreaterEqual(len(samples), 0) if len(samples) > 0: # Sample should be truncated to song length sample_duration = samples[0]["data"].shape[1] / self.cfg.semantic_rate_hz self.assertLessEqual(sample_duration, 4.0) def test_song_to_samples_no_sample_idx(self): """Test song_to_samples when no sample_idx is provided.""" meta_no_idx = {k: v for k, v in self.song_meta.items() if k != "sample_idx"} samples = song_to_samples(self.song_data, meta_no_idx, self.cfg) self.assertEqual(len(samples), 1) # Should still create a sample, just without sample_idx reference self.assertIsInstance(samples[0]["data"], np.ndarray) def test_get_samples_for_song(self): """Test get_samples_for_song helper function.""" # Mock data and metas data = [self.song_data, np.random.randint(0, 4096, (1, 500))] metas = [self.song_meta, {"duration": 20.0, "sample_idx": 789}] samples = get_samples_for_song(0, data, metas, self.cfg) self.assertIsInstance(samples, list) self.assertGreater(len(samples), 0) for sample in samples: self.assertIn("data", sample) self.assertIn("meta", sample) def test_sample_block_type(self): """Test SampleBlockType properties.""" self.assertEqual(SampleBlockType.name, "sample") self.assertTrue(SampleBlockType.is_causal) def test_gpt_config_semantic_sample_token(self): """Test GPTConfig semantic_sample_token initialization.""" # Use realistic values that match the default config config = GPTConfig(semantic_codebook_size=4000, semantic_vocab_size=4032) # Should auto-initialize to codebook_size + 12 expected_token = 4000 + 12 self.assertEqual(config.semantic_sample_token, expected_token) # Test explicit setting config_explicit = GPTConfig( semantic_codebook_size=4000, semantic_vocab_size=4032, semantic_sample_token=4020 ) self.assertEqual(config_explicit.semantic_sample_token, 4020) def test_sampling_params_defaults(self): """Test SamplingParams default values.""" params = SamplingParams() self.assertFalse(params.allow_sample) self.assertEqual(params.prob_sample, 0.1) def test_sampling_params_custom(self): """Test SamplingParams with custom values.""" params = SamplingParams(allow_sample=True, prob_sample=0.2) self.assertTrue(params.allow_sample) self.assertEqual(params.prob_sample, 0.2) def test_sample_data_structure(self): """Test SampleData structure and properties.""" sample_data = SampleData( data_row=self.song_data, data_meta=self.song_meta, sampling_params=self.sampling_params, audio_sample_tracks=self.sample_data.audio_sample_tracks, ) self.assertIsNotNone(sample_data.audio_sample_tracks) self.assertEqual(sample_data.audio_sample_tracks[0].shape, (1, 250)) def test_sample_duration_validation(self): """Test sample duration constraints.""" # Test various sample durations durations = [5, 10, 15, 20] # seconds for duration in durations: sample_tokens = int(duration * self.cfg.semantic_rate_hz) sample_data = np.random.randint(0, 4096, (1, sample_tokens)) # Verify token count matches expected duration calculated_duration = sample_data.shape[1] / self.cfg.semantic_rate_hz self.assertAlmostEqual(calculated_duration, duration, places=1) def test_semantic_sample_token_validation(self): """Test semantic_sample_token is within vocabulary bounds.""" config = GPTConfig(semantic_codebook_size=4000, semantic_vocab_size=4032) # Token should be less than vocab size self.assertLess(config.semantic_sample_token, config.semantic_vocab_size) # Should be greater than codebook size (reserved range) self.assertGreater(config.semantic_sample_token, config.semantic_codebook_size) def test_backward_compatibility(self): """Test that new features don't break existing functionality.""" # Test SamplingParams with allow_sample=False (default) params = SamplingParams(allow_sample=False) self.assertFalse(params.allow_sample) # Test SampleData without audio_sample_tracks (but with required params) sample_data = SampleData( data_row=self.song_data, data_meta=self.song_meta, sampling_params=params ) self.assertIsNone(sample_data.audio_sample_tracks) # These should not cause errors in existing code paths def test_sample_metadata_preservation(self): """Test that sample metadata is correctly preserved.""" samples = song_to_samples(self.song_data, self.song_meta, self.cfg, n_samples=1) sample = samples[0] self.assertIn("meta", sample) # Original metadata should be preserved in sample original_keys = ["duration", "sample_idx", "artist", "tags"] for key in original_keys: if key in self.song_meta: # Sample meta should reference or contain original info self.assertIsInstance(sample["meta"], dict) def test_integration_sample_to_song_workflow(self): """Test the complete sample-to-song workflow integration.""" # Step 1: Create samples from song samples = song_to_samples(self.song_data, self.song_meta, self.cfg, n_samples=2) # Step 2: Verify samples can be used for conditioning for sample in samples: self.assertIsInstance(sample["data"], np.ndarray) self.assertEqual(sample["data"].shape[0], 1) # Single codebook # Sample should be suitable for SampleData sample_data = SampleData( data_row=self.song_data, data_meta=self.song_meta, sampling_params=self.sampling_params, audio_sample_tracks=[sample["data"]], ) self.assertIsNotNone(sample_data.audio_sample_tracks) # Verify duration constraints duration = sample["data"].shape[1] / self.cfg.semantic_rate_hz self.assertGreaterEqual(duration, 5) self.assertLessEqual(duration, 15) def test_error_handling(self): """Test error handling in sample processing.""" # Test with missing required config attributes incomplete_cfg = Mock() incomplete_cfg.semantic_rate_hz = "invalid" # String instead of int # This should cause a TypeError when multiplying with self.assertRaises((TypeError, AttributeError)): song_to_samples(self.song_data, self.song_meta, incomplete_cfg) # Test with None data - this should handle gracefully or raise appropriate error try: result = song_to_samples(None, self.song_meta, self.cfg) # If it doesn't raise an error, it should return empty list or similar self.assertIsInstance(result, list) except (ValueError, TypeError, AttributeError): # These are acceptable error types for None input pass def test_sample_timing_text_functionality(self): """Test that sample timing is correctly added to text conditioning.""" from text_utils import build_text # Test case: sample at 1:43.4 (103.4 seconds) result = build_text( tags=["pop", "energetic"], text="Test lyrics here", sample_duration_s=180, sample_duration_toks=4500, inference=True, # Use inference mode for consistent output audio_sample_start_times_s=[103.4], ) # Should contain the timing tag in control tags format (seconds) self.assertIn("audio_sample_time_0:103", result) # Should also contain token-based timing self.assertIn("audio_sample_start_toks_0:2585", result) # 103.4 * 25 ≈ 2585 # Test without sample timing result_no_timing = build_text( tags=["test"], text="Test lyrics", sample_duration_s=180, sample_duration_toks=4500, inference=True, audio_sample_start_times_s=None, ) # Should not contain audio sample timing tag self.assertNotIn("audio_sample_time_", result_no_timing) # Test different timing values with zero-based indexing test_cases = [ (12.1, "audio_sample_time_0:12"), (176.6, "audio_sample_time_0:177"), (31.2, "audio_sample_time_0:31"), (3661.5, "audio_sample_time_0:3662"), # Over 1 hour ] for start_time, expected_tag in test_cases: result = build_text( tags=["test"], text="lyrics", sample_duration_s=30, sample_duration_toks=750, inference=True, audio_sample_start_times_s=[start_time], ) self.assertIn(expected_tag, result, f"Failed for timing {start_time}s") def test_multiple_audio_sample_timing(self): """Test multiple audio sample timings in control tags.""" from text_utils import build_text # Test case: multiple samples at different times result = build_text( tags=["electronic", "energetic"], text="Multiple samples test", sample_duration_s=180, sample_duration_toks=4500, inference=True, audio_sample_start_times_s=[12.3, 67.8, 134.5], ) # Should contain multiple timing tags in seconds format with zero-based indexing self.assertIn("audio_sample_time_0:12", result) self.assertIn("audio_sample_time_1:68", result) self.assertIn("audio_sample_time_2:134", result) # 134.5 rounds to 134 # Should also contain token-based timing (25 Hz semantic rate) self.assertIn("audio_sample_start_toks_0:308", result) # 12.3 * 25 ≈ 308 self.assertIn("audio_sample_start_toks_1:1695", result) # 67.8 * 25 ≈ 1695 self.assertIn("audio_sample_start_toks_2:3362", result) # 134.5 * 25 ≈ 3362 def test_randomized_sample_count_range(self): """Test that max_num_audio_samples creates variable sample counts.""" import random from data_types import SamplingParams # Test the randomization logic max_samples = 4 counts = [] random.seed(42) # For reproducible test for _ in range(20): # Test multiple iterations count = random.randint(1, max_samples) counts.append(count) # Should have variety in counts unique_counts = set(counts) self.assertGreater(len(unique_counts), 1, "Should generate different sample counts") self.assertGreaterEqual(min(counts), 1, "Minimum should be 1") self.assertLessEqual(max(counts), max_samples, f"Maximum should be {max_samples}") # Test SamplingParams structure params = SamplingParams(max_num_audio_samples=10) self.assertEqual(params.max_num_audio_samples, 10) def test_audio_sample_source_control_tags(self): """Test audio sample source control tags generation.""" from text_utils import build_text # Test single sample with source result_single = build_text( tags=["electronic"], text="Single source test", sample_duration_s=30, sample_duration_toks=750, inference=True, audio_sample_start_times_s=[45.2], audio_sample_sources=["vocal"], ) # Should contain timing and source tags with zero-based indexing self.assertIn("audio_sample_time_0:45", result_single) self.assertIn("audio_sample_start_toks_0:1130", result_single) # 45.2 * 25 ≈ 1130 self.assertIn("audio_sample_vocal_0", result_single) # Test multiple samples with different sources result_multiple = build_text( tags=["rock"], text="Multiple sources test", sample_duration_s=120, sample_duration_toks=3000, inference=True, audio_sample_start_times_s=[12.1, 67.8, 134.5], audio_sample_sources=["vocal", "drum", "full_mix"], ) # Should contain timing tags self.assertIn("audio_sample_time_0:12", result_multiple) self.assertIn("audio_sample_time_1:68", result_multiple) self.assertIn("audio_sample_time_2:134", result_multiple) # 134.5 rounds to 134 # Should contain source tags self.assertIn("audio_sample_vocal_0", result_multiple) self.assertIn("audio_sample_drum_1", result_multiple) self.assertIn("audio_sample_full_mix_2", result_multiple) # Test without sources (backward compatibility) result_no_sources = build_text( tags=["jazz"], text="No sources test", sample_duration_s=60, sample_duration_toks=1500, inference=True, audio_sample_start_times_s=[23.4], audio_sample_sources=None, ) # Should contain timing but no source tags self.assertIn("audio_sample_time_0:23", result_no_sources) self.assertNotIn("audio_sample_vocal", result_no_sources) self.assertNotIn("audio_sample_drum", result_no_sources) self.assertNotIn("audio_sample_full_mix", result_no_sources) if __name__ == "__main__": unittest.main()