diff --git a/breads/instruments/KPIC.py b/breads/instruments/KPIC.py index 2ecd09b..f004b1a 100644 --- a/breads/instruments/KPIC.py +++ b/breads/instruments/KPIC.py @@ -298,22 +298,65 @@ def get_fib_labels(header): # else: # return np.array(sf_id_list)[np.array(sf_num_list,dtype=np.float).argsort()] -def prep_data_object(date, file_numbers, fiber_list, datadir, trace_filename, wvs_filename, orders, bkgdsub=True): + +def prepare_all_data(date, file_numbers, fiber_list, datadir, trace_filename, + wvs_filename, orders, bkgdsub=True): """ - Prepares a KPIC data object per fiber, returning a dictionary keyed by fiber number (0–3). - Only fibers with associated files are included. + Prepares KPIC data objects segregated by science fiber number. + + This function identifies all discrete science fiber numbers from the fiber_list, + groups the corresponding file numbers, and creates separate data objects for each fiber. + Science fibers are indexed from 0. + + Parameters + ---------- + date : str + Date string for constructing filenames + file_numbers : array-like + Array of file numbers corresponding to each observation + fiber_list : array-like + Array of fiber numbers (0-3) corresponding to each file in file_numbers. + Must be same length as file_numbers. + datadir : str + Directory path containing the data files + trace_filename : str + Path to trace calibration file + wvs_filename : str + Path to wavelength calibration file + orders : array-like + Spectral orders to select + bkgdsub : bool, optional + If True, use background-subtracted files (bkgdsub), else use nodding subtraction (nodsub). + Default is True. + + Returns + ------- + dict + Dictionary of KPIC data objects keyed by fiber number (0–3). + Only fibers with associated files are included. + + Examples + -------- + >>> # Example with files for fibers 0 and 2 + >>> fiber_list = [0, 0, 2, 2, 0] + >>> file_numbers = [100, 101, 102, 103, 104] + >>> data_objects = prepare_all_data('20210101', file_numbers, fiber_list, + ... '/data/', 'trace.fits', 'wvs.fits', [5, 6, 7]) + >>> # Returns: {0: , 2: } """ fiber_dataobjs = {} fiber_list = np.array(fiber_list) file_numbers = np.array(file_numbers) - for fiber in range(4): + # Identify all discrete science fiber numbers present in the data + unique_fibers = np.unique(fiber_list) + + for fiber in unique_fibers: + # Find all indices corresponding to this science fiber indices = np.where(fiber_list == fiber)[0] - if len(indices) == 0: - continue - print(f"Fiber {fiber} has {len(indices)} files") + print(f"Science fiber {fiber} has {len(indices)} files") filelist = [] for idx in indices: filenum = file_numbers[idx] @@ -321,14 +364,30 @@ def prep_data_object(date, file_numbers, fiber_list, datadir, trace_filename, wv filepath = os.path.join(datadir, filename) filelist.append(filepath) + # Create data object for this fiber with all its associated files fiber_goal_list = [fiber] * len(filelist) - dataobj = KPIC(filelist, trace_filename, wvs_filename, combine_mode="companion", fiber_goal_list=fiber_goal_list) + dataobj = KPIC(filelist, trace_filename, wvs_filename, + combine_mode="companion", fiber_goal_list=fiber_goal_list) dataobj = dataobj.selec_order(orders) fiber_dataobjs[fiber] = dataobj return fiber_dataobjs -def prep_host_object(date, file_numbers, fiber_list, datadir, trace_filename, wvs_filename, orders, bkgdsub = True): + +# Backward compatibility alias +def prep_data_object(date, file_numbers, fiber_list, datadir, trace_filename, + wvs_filename, orders, bkgdsub=True): + """ + Deprecated: Use prepare_all_data() instead. + + This function is kept for backward compatibility and simply calls prepare_all_data(). + """ + return prepare_all_data(date, file_numbers, fiber_list, datadir, + trace_filename, wvs_filename, orders, bkgdsub) + + +def prep_host_object(date, file_numbers, fiber_list, datadir, trace_filename, + wvs_filename, orders, bkgdsub=True): """ Prepares a data object for the given file numbers and fiber. """ @@ -338,11 +397,13 @@ def prep_host_object(date, file_numbers, fiber_list, datadir, trace_filename, wv filelist.append(os.path.join(datadir, f"nspec{date}_{filenum:04d}_bkgdsub_spectra.fits")) else: filelist.append(os.path.join(datadir, f"nspec{date}_{filenum:04d}_nodsub_spectra.fits")) - + dataobj = KPIC(filelist, trace_filename, wvs_filename, combine_mode="star", fiber_goal_list=fiber_list) return dataobj.selec_order(orders) -def prep_A0_object(date, file_numbers, fiber_list, datadir, trace_filename, wvs_filename, orders, bkgdsub = True): + +def prep_A0_object(date, file_numbers, fiber_list, datadir, trace_filename, + wvs_filename, orders, bkgdsub=True): """ Prepares a data object for the given file numbers and fiber. """ @@ -352,6 +413,6 @@ def prep_A0_object(date, file_numbers, fiber_list, datadir, trace_filename, wvs_ filelist.append(os.path.join(datadir, f"nspec{date}_{filenum:04d}_bkgdsub_spectra.fits")) else: filelist.append(os.path.join(datadir, f"nspec{date}_{filenum:04d}_nodsub_spectra.fits")) - + dataobj = KPIC(filelist, trace_filename, wvs_filename, combine_mode="star", fiber_goal_list=fiber_list) - return dataobj.selec_order(orders) \ No newline at end of file + return dataobj.selec_order(orders) diff --git a/breads/tests/test_kpic.py b/breads/tests/test_kpic.py new file mode 100644 index 0000000..22ec22d --- /dev/null +++ b/breads/tests/test_kpic.py @@ -0,0 +1,106 @@ +""" +Tests for KPIC instrument module, specifically the prepare_all_data function. +""" +import numpy as np + + +def test_prepare_all_data_import(): + """Test that prepare_all_data can be imported from KPIC module""" + from breads.instruments.KPIC import prepare_all_data + assert callable(prepare_all_data) + + +def test_prep_data_object_backward_compatibility(): + """Test that old function name still works for backward compatibility""" + from breads.instruments.KPIC import prep_data_object + assert callable(prep_data_object) + + +def test_fiber_segregation_logic(): + """ + Test the logic of segregating files by fiber number. + This tests the core logic without requiring actual data files. + """ + # Test case 1: Multiple fibers with multiple files each + fiber_list = np.array([0, 0, 2, 2, 0, 1, 1]) + file_numbers = np.array([100, 101, 102, 103, 104, 105, 106]) + + # Expected groupings: + # Fiber 0: files 100, 101, 104 + # Fiber 1: files 105, 106 + # Fiber 2: files 102, 103 + + # Check unique fibers + unique_fibers = np.unique(fiber_list) + assert len(unique_fibers) == 3 + assert 0 in unique_fibers + assert 1 in unique_fibers + assert 2 in unique_fibers + + # Check fiber 0 indices + indices_fiber_0 = np.where(fiber_list == 0)[0] + assert len(indices_fiber_0) == 3 + assert np.array_equal(file_numbers[indices_fiber_0], [100, 101, 104]) + + # Check fiber 1 indices + indices_fiber_1 = np.where(fiber_list == 1)[0] + assert len(indices_fiber_1) == 2 + assert np.array_equal(file_numbers[indices_fiber_1], [105, 106]) + + # Check fiber 2 indices + indices_fiber_2 = np.where(fiber_list == 2)[0] + assert len(indices_fiber_2) == 2 + assert np.array_equal(file_numbers[indices_fiber_2], [102, 103]) + + +def test_fiber_indexing_from_zero(): + """Test that science fibers are indeed indexed from 0""" + fiber_list = np.array([0, 1, 2, 3]) + unique_fibers = np.unique(fiber_list) + + # Verify all fiber indices are non-negative (indexed from 0) + assert np.all(unique_fibers >= 0) + assert 0 in unique_fibers + + +def test_empty_fiber_handling(): + """Test that missing fibers are handled correctly""" + # Only fibers 0 and 2 present, fibers 1 and 3 missing + fiber_list = np.array([0, 0, 2, 2]) + + unique_fibers = np.unique(fiber_list) + + # Should only have fibers 0 and 2 + assert len(unique_fibers) == 2 + assert 0 in unique_fibers + assert 2 in unique_fibers + assert 1 not in unique_fibers + assert 3 not in unique_fibers + + +def test_single_fiber_single_file(): + """Test edge case with single fiber and single file""" + fiber_list = np.array([0]) + file_numbers = np.array([100]) + + unique_fibers = np.unique(fiber_list) + assert len(unique_fibers) == 1 + assert unique_fibers[0] == 0 + + indices = np.where(fiber_list == 0)[0] + assert len(indices) == 1 + assert file_numbers[indices[0]] == 100 + + +def test_all_files_same_fiber(): + """Test when all files belong to the same fiber""" + fiber_list = np.array([2, 2, 2, 2]) + file_numbers = np.array([100, 101, 102, 103]) + + unique_fibers = np.unique(fiber_list) + assert len(unique_fibers) == 1 + assert unique_fibers[0] == 2 + + indices = np.where(fiber_list == 2)[0] + assert len(indices) == 4 + assert np.array_equal(file_numbers[indices], [100, 101, 102, 103])