#!/usr/bin/env python3 """Test script for merge_preference_datasets. This script creates minimal test datasets and validates the merge operation. """ import os import tempfile import shutil import numpy as np from typing import Dict, Any, List import json from merge_preference_datasets import ( merge_preference_datasets, load_sample_from_mmap, read_jsonl, read_json, write_jsonl, write_json, SEMANTIC_N_CODEBOOKS, N_TOKENS_AUDIO, ) def create_test_dataset( output_dir: str, n_samples: int, is_val: bool = False, t_data_memmap: int = 100, # Smaller for testing start_id: int = 0, ) -> None: """Create a minimal test dataset. Args: output_dir: Directory to create the test dataset n_samples: Number of samples to create is_val: Whether this is validation set t_data_memmap: Number of tokens per sample (smaller for tests) start_id: Starting ID for sample naming """ os.makedirs(output_dir, exist_ok=True) dset_type = "val" if is_val else "tr" mmap_path = os.path.join(output_dir, f"data_{dset_type}.bin") meta_path = os.path.join(output_dir, f"meta_{dset_type}.jsonl") info_path = os.path.join(output_dir, f"info_{dset_type}.json") # Create mmap sample_size = t_data_memmap * SEMANTIC_N_CODEBOOKS total_size = n_samples * sample_size mmap = np.memmap(mmap_path, dtype=np.uint16, mode='w+', shape=(total_size,)) # Create metadata and info metadata: List[Dict[str, Any]] = [] info: Dict[str, Dict[str, Any]] = { "perference_0": {"idx_list": []}, "perference_1": {"idx_list": []}, } for i in range(n_samples): # Create sample data with unique pattern sample_data = np.full( (t_data_memmap, SEMANTIC_N_CODEBOOKS), fill_value=i + start_id, dtype=np.uint16 ) # Write to mmap offset = i * sample_size mmap[offset:offset + sample_size] = sample_data.flatten() # Create metadata preference = i % 2 meta = { "dataset": f"perference_{preference}", "id": f"test_sample_{i + start_id}", "start_s": 0.0, "vocal_start_s": None, "vocal_end_s": None, "tags": ["test"], "neg_tags": "", "control_tags": "", "gender": "male" if i % 2 == 0 else "female", "control": {}, "text": f"Test sample {i + start_id}", "generated_start_index": 0, "user_id": f"user_{i % 3}", } metadata.append(meta) info[f"perference_{preference}"]["idx_list"].append(i) mmap.flush() del mmap # Write metadata and info write_jsonl(metadata, meta_path) write_json(info, info_path) print(f"Created test dataset at {output_dir} with {n_samples} samples") def test_merge() -> None: """Test the merge_preference_datasets function.""" print("=" * 70) print("TESTING MERGE_PREFERENCE_DATASETS") print("=" * 70) # Create temporary directories temp_dir = tempfile.mkdtemp(prefix="test_merge_") try: # Test parameters t_data_memmap = 100 # Small for fast testing n_samples1 = 20 n_samples2 = 30 input_dir1 = os.path.join(temp_dir, "dataset1") input_dir2 = os.path.join(temp_dir, "dataset2") output_dir = os.path.join(temp_dir, "merged") print(f"\nTest parameters:") print(f" t_data_memmap: {t_data_memmap}") print(f" Dataset 1 samples: {n_samples1}") print(f" Dataset 2 samples: {n_samples2}") print(f" Temp directory: {temp_dir}") # Create test datasets print("\n" + "=" * 70) print("Creating test datasets...") print("=" * 70) create_test_dataset(input_dir1, n_samples1, is_val=True, t_data_memmap=t_data_memmap, start_id=0) create_test_dataset(input_dir2, n_samples2, is_val=True, t_data_memmap=t_data_memmap, start_id=n_samples1) # Merge datasets print("\n" + "=" * 70) print("Merging datasets...") print("=" * 70) merge_preference_datasets( input_dir1=input_dir1, input_dir2=input_dir2, output_dir=output_dir, is_val=True, t_data_memmap=t_data_memmap, validate=True, ) # Verify merged dataset print("\n" + "=" * 70) print("Verifying merged dataset...") print("=" * 70) merged_mmap_path = os.path.join(output_dir, "data_val.bin") merged_meta_path = os.path.join(output_dir, "meta_val.jsonl") merged_info_path = os.path.join(output_dir, "info_val.json") # Load merged data merged_meta = read_jsonl(merged_meta_path) merged_info = read_json(merged_info_path) # Check counts assert len(merged_meta) == n_samples1 + n_samples2, \ f"Metadata count mismatch: {len(merged_meta)} != {n_samples1 + n_samples2}" print(f"✓ Metadata count correct: {len(merged_meta)}") # Check info indices total_indices = sum(len(v["idx_list"]) for v in merged_info.values()) assert total_indices == n_samples1 + n_samples2, \ f"Info indices count mismatch: {total_indices} != {n_samples1 + n_samples2}" print(f"✓ Info indices count correct: {total_indices}") # Verify samples from dataset 1 print("\nVerifying samples from dataset 1...") for i in range(min(5, n_samples1)): sample = load_sample_from_mmap(merged_mmap_path, i, t_data_memmap) expected_value = i actual_value = sample[0, 0] assert actual_value == expected_value, \ f"Dataset 1 sample {i} mismatch: {actual_value} != {expected_value}" meta = merged_meta[i] assert meta["id"] == f"test_sample_{i}", \ f"Dataset 1 metadata {i} ID mismatch" print(f"✓ Dataset 1 samples verified (checked {min(5, n_samples1)} samples)") # Verify samples from dataset 2 print("\nVerifying samples from dataset 2...") for i in range(min(5, n_samples2)): merged_idx = n_samples1 + i sample = load_sample_from_mmap(merged_mmap_path, merged_idx, t_data_memmap) expected_value = n_samples1 + i actual_value = sample[0, 0] assert actual_value == expected_value, \ f"Dataset 2 sample {i} mismatch: {actual_value} != {expected_value}" meta = merged_meta[merged_idx] assert meta["id"] == f"test_sample_{n_samples1 + i}", \ f"Dataset 2 metadata {i} ID mismatch" print(f"✓ Dataset 2 samples verified (checked {min(5, n_samples2)} samples)") # Verify info indices are correctly offset print("\nVerifying info indices...") for dataset_name, dataset_info in merged_info.items(): idx_list = dataset_info["idx_list"] for idx in idx_list: meta = merged_meta[idx] assert meta["dataset"] == dataset_name, \ f"Info index {idx} points to wrong dataset: {meta['dataset']} != {dataset_name}" print("✓ All info indices point to correct metadata") print("\n" + "=" * 70) print("✅ ALL TESTS PASSED!") print("=" * 70) finally: # Clean up if os.path.exists(temp_dir): shutil.rmtree(temp_dir) print(f"\nCleaned up temp directory: {temp_dir}") if __name__ == "__main__": test_merge()