Commit a63e6632 authored by Fang Yuedong's avatar Fang Yuedong
Browse files

add first measurement and evaluation pipelines

parent 33baa6e6
Loading
Loading
Loading
Loading
+18 −4
Original line number Diff line number Diff line
@@ -8,16 +8,16 @@
# n_objects: 500
rotate_objs: NO
use_mpi: YES
run_name: "test_20230509"
run_name: "test_20230517"

pos_sampling:
  # type: "HexGrid"
  # type: "RectGrid"
  # grid_spacing: 18.5 # arcsec (~250 pixels)
  type: "uniform"
  object_density: 50 # arcmin^-2
  object_density: 37 # arcmin^-2

output_img_dir: "./workspace"
output_img_dir: "/share/home/fangyuedong/injection_pipeline/workspace"
input_img_list: "/share/home/fangyuedong/injection_pipeline/input_L1_img_MSC_0000000.list"

###############################################
@@ -80,3 +80,17 @@ ins_effects:
###############################################
random_seeds:
  seed_Av:              121212    # Seed for generating random intrinsic extinction

###############################################
# Measurement setting
###############################################
measurement_setting:
  input_img_list: "/share/home/fangyuedong/injection_pipeline/injected_L1_img_MSC_0000000.list"
  # input_img_list: "/share/home/fangyuedong/injection_pipeline/input_L1_img_MSC_0000000.list"
  input_wht_list: "/share/home/fangyuedong/injection_pipeline/L1_wht_img_MSC_0000000.list"
  input_flg_list: "/share/home/fangyuedong/injection_pipeline/L1_flg_img_MSC_0000000.list"
  input_psf_list: "/share/home/fangyuedong/injection_pipeline/psf_img_MSC_0000000.list"
  sex_config: "/share/home/fangyuedong/injection_pipeline/config/default.config"
  sex_param: "/share/home/fangyuedong/injection_pipeline/config/default.param"
  n_jobs: 8
  output_dir: "/share/home/fangyuedong/injection_pipeline/workspace"
 No newline at end of file

evaluation/__init__.py

0 → 100644
+0 −0

Empty file added.

+174 −0
Original line number Diff line number Diff line
import argparse
import numpy as np
import os
from astropy.io import fits
from astropy.stats import sigma_clip
from astropy.stats import median_absolute_deviation
from astropy.stats import mad_std

def get_all_stats(values_arrays):
    """
    """
    stats = {}

    stats['mean']   = np.mean(values_arrays)
    stats['median'] = np.median(values_arrays)
    stats['std']    = np.std(values_arrays)
    stats['mad']    = mad_std(values_arrays)

    return stats

def define_options():
    parser = argparse.ArgumentParser()
    parser.add_argument('--data_image', dest='data_image', type=str, required=True,
                        help='Name of the data image: (default: "%(default)s"')
    parser.add_argument('--seg_image', dest='seg_image', type=str, required=True,
                        help='Name of the mask / segmentation image: (default: "%(default)s"')
    parser.add_argument('--flag_image', dest='flag_image', type=str, required=False,
                        default=None, help='Name of the flag image (default: "%(default)s"')
    parser.add_argument('--sky_image', dest='sky_image', type=str, required=False,
                        default=None, help='Name of the sky image (default: "%(default)s"')
    parser.add_argument('--aper_min', dest='aper_min', type=int, required=False,
                        default=5, help='Minimum no. of pixels at level: (default: "%(default)s"')
    parser.add_argument('--aper_max', dest='aper_max', type=int, required=False,
                        default=20, help='Maximum no. of pixels at level: (default: "%(default)s"')
    parser.add_argument('--aper_sampling', dest='aper_sampling', type=int, required=False,
                        default=1, help='Minimum no. of pixels at level: (default: "%(default)s"')
    parser.add_argument('--n_sample', dest='n_sample', type=int, required=False,
                        default=500, help='Minimum no. of pixels at level: (default: "%(default)s"')
    parser.add_argument('--out_basename', dest='out_basename', type=str, required=False,
                        default="aper", help='Base name for the output names: (default: "%(default)s"')
    parser.add_argument('--output_dir', dest='output_dir', type=str, required=False,
                        default="./workspace", help='dir path for the output : (default: "%(default)s"')
    return parser

def create_circular_mask(h, w, center=None, radius=None):
    if center is None: # use the middle of the image
        center = [int(w / 2), int(h / 2)]
    if radius is None: # use the smallest distance between the center and image walls
        radius = min(center[0], center[1], w - center[0], h - center[1])

    Y, X = np.ogrid[:h, :w]
    dist_from_center = np.sqrt((X - center[0]) ** 2 + (Y - center[1]) ** 2)

    mask = dist_from_center <= radius
    return mask

def sampling(image_data, seg_data, flag_data, aperture, Nsample):
    wx, wy = np.where( (image_data != 0) & (seg_data == 0) & (flag_data == 0))
    Nx, Ny = image_data.shape

    flux_average = np.zeros(Nsample)
    flux_median = np.zeros(Nsample)
    x_position = np.zeros(Nsample)
    y_position = np.zeros(Nsample)
    i = 0
    i_iter = 0
    while i < Nsample:
        if i_iter > 100*Nsample:
            print('# Not enough background pixels for image depth analysis!')
            break
        i_iter += 1

        idx = np.random.randint(len(wx))
        stmpsize = aperture+1

        if wx[idx]+stmpsize >= Nx:
            continue
        if wy[idx]+stmpsize >= Ny:
            continue

        img_stmp  = image_data[ wx[idx]:wx[idx]+stmpsize, wy[idx]:wy[idx]+stmpsize ]
        seg_stmp  = seg_data[ wx[idx]:wx[idx]+stmpsize, wy[idx]:wy[idx]+stmpsize ]
        flag_stmp = flag_data[ wx[idx]:wx[idx]+stmpsize, wy[idx]:wy[idx]+stmpsize ]

        mask = create_circular_mask(stmpsize, stmpsize, center=[stmpsize//2,stmpsize//2], radius=aperture//2)
        area = np.pi*(aperture/2)**2
        area_sum = len(mask[mask==True])
        ratio = area/area_sum

        ss = np.sum(seg_stmp[mask])
        if ss != 0:
            continue
        fs = np.sum(flag_stmp[mask])
        if fs != 0:
            continue
        flux_average[i] = np.average(img_stmp[mask])
        flux_median[i]  = np.median(img_stmp[mask])
        x_position[i] = (wx[idx]+wx[idx]+stmpsize)/2.0
        y_position[i] = (wy[idx]+wy[idx]+stmpsize)/2.0 

        i += 1

    print('Needed %i tries for %i samples!'%(i_iter, Nsample))

    return flux_average, flux_median, x_position, y_position
    
def noise_statistics_aperture(fitsname, segname, flagname=None, sky_image=None, aperture_min=1, aperture_max=10, aperture_step=1, seed=None, Nsample=100, sigma_cl=10., base_name="aper", output_dir='./'):
    f     = fits.open(fitsname)
    fseg  = fits.open(segname)
    # image_data = f[1].data
    image_data = f[1].data * f[1].header["GAIN1"]
    seg_data = fseg[0].data

    f.close()
    fseg.close()

    if flagname:
        fflag = fits.open(flagname)
        flag_data = fflag[1].data
        fflag.close()
    else:
        flag_data = np.zeros(image_data.shape)

    if sky_image:
        hdu = fits.open(sky_image)
        sky_data = hdu[0].data * hdu[0].header["GAIN1"]
        image_data -= sky_data

    if seed != None:
        np.random.seed(seed)
    
    if not os.path.exists(output_dir):
        os.makedirs(output_dir)

    im = image_data
    im[seg_data > 0] = 0.
    hist_data = im[im != 0.].flatten()

    aperture_list = np.arange(aperture_min, aperture_max+1, aperture_step, dtype=int)
    sigma_output = np.zeros(len(aperture_list))
    mad_output   = np.zeros(len(aperture_list))
    mad_std_output = np.zeros(len(aperture_list))

    for j, aperture in enumerate(aperture_list):
        flux_average, flux_median, x_position, y_position = sampling(image_data, seg_data, flag_data, aperture, Nsample)
        
        mean_stats   = get_all_stats(flux_average)
        median_stats = get_all_stats(flux_median)
        print("Mean:   %e += %e +- %e"%(mean_stats['median'], mean_stats['mad'], mean_stats['std']))
        print("Median: %e += %e +- %e"%(median_stats['median'], median_stats['mad'], median_stats['std']))

        aper_file = '%s_%03i.txt'%(base_name, aperture)
        aper_file = os.path.join(output_dir, aper_file)
        print('Aperture file: %s'%aper_file)
        with open(aper_file, "w+") as aper_out:
            for one_value in zip(flux_average, flux_median, x_position, y_position):
                one_line = "{:.7f} {:.7f} {:.1f} {:.1f}\n".format(*one_value)
                aper_out.write(one_line)

    return aperture_list, sigma_output, mad_output, mad_std_output

if __name__ == "__main__":
    args = define_options().parse_args()
    aperture_ap, sigma_ap, mad_ap, nmad_ap = noise_statistics_aperture(
        args.data_image,
        args.seg_image,
        args.flag_image,
        args.sky_image,
        aperture_min=args.aper_min,
        aperture_max=args.aper_max,
        aperture_step=args.aper_sampling,
        Nsample=args.n_sample,
        base_name=args.out_basename,
        output_dir=args.output_dir)
 No newline at end of file
+184 −0
Original line number Diff line number Diff line
import argparse
import os
import numpy as np
import matplotlib.pyplot as plt
from astropy.io import ascii, fits
from cross_match_catalogs import read_catalog, match_catalogs_img
from ObservationSim.Instrument import Telescope, Filter, FilterParam

VC_A = 2.99792458e+18  # speed of light: A/s

def define_options():
    parser = argparse.ArgumentParser()
    parser.add_argument('--TU_catalog', dest='TU_catalog', type=str, required=True,
                        help='path to the (injected) truth catalog')
    parser.add_argument('--source_catalog', dest='source_catalog', type=str, required=True,
                        help='path to the (extracted) injected catalog')
    parser.add_argument('--orig_catalog', dest='orig_catalog', type=str, required=True,
                        help='path to the (extracted) original catalog')
    parser.add_argument('--image', dest='image', type=str, required=True,
                        help='path to the image, used to get the header info')
    parser.add_argument('--output_dir', dest='output_dir', type=str, required=False,
                        default='./workspace', help='output path')
    return parser

def getChipFilter(chipID, filter_layout=None):
        """Return the filter index and type for a given chip #(chipID)
        """
        filter_type_list = ["nuv","u", "g", "r", "i","z","y","GU", "GV", "GI", "FGS"]
        if filter_layout is not None:
            return filter_layout[chipID][0], filter_layout[chipID][1]

        # updated configurations
        if chipID>42 or chipID<1: raise ValueError("!!! Chip ID: [1,42]")
        if chipID in [6, 15, 16, 25]:  filter_type = "y"
        if chipID in [11, 20]:         filter_type = "z"
        if chipID in [7, 24]:          filter_type = "i"
        if chipID in [14, 17]:         filter_type = "u"
        if chipID in [9, 22]:          filter_type = "r"
        if chipID in [12, 13, 18, 19]: filter_type = "nuv"
        if chipID in [8, 23]:          filter_type = "g"
        if chipID in [1, 10, 21, 30]:  filter_type = "GI"
        if chipID in [2, 5, 26, 29]:   filter_type = "GV"
        if chipID in [3, 4, 27, 28]:   filter_type = "GU"
        if chipID in range(31, 43):    filter_type = 'FGS'
        filter_id = filter_type_list.index(filter_type)

        return filter_id, filter_type

def magToFlux(mag):
    """
    flux of a given AB magnitude

    Parameters:
    mag: magnitude in unit of AB

    Return:
    flux: flux in unit of erg/s/cm^2/Hz
    """
    flux = 10**(-0.4*(mag+48.6))
    return flux

def getElectronFluxFilt(mag, filt, tel, exptime=150.):
    photonEnergy = filt.getPhotonE()
    flux = magToFlux(mag)
    factor = 1.0e4 * flux/photonEnergy * VC_A * (1.0/filt.blue_limit - 1.0/filt.red_limit)
    return factor * filt.efficiency * tel.pupil_area * exptime

def convert_catalog(catname):
    data_dir = os.path.dirname(catname)
    base_name = os.path.basename(catname)
    text_file = ascii.read(catname)
    fits_filename = os.path.join(data_dir, base_name + '.fits')
    text_file.write(fits_filename, overwrite=True)

def validation_hist(val, idx, name="val", nbins=10, fig_name='detected_counts.png', output_dir='./'):
    counts, bins = np.histogram(val, bins=nbins)
    is_empty = np.full(len(val), False)
    for i in range(len(idx)):
        if idx[i].size == 0:
            is_empty[i] = True
    counts_detected, _ = np.histogram(val[~is_empty], bins=nbins)
    plt.figure()
    plt.stairs(counts, bins, color='r', label='TU objects')
    plt.stairs(counts_detected, bins, color='g', label='Detected')
    plt.xlabel(name, size='x-large')
    plt.title("Counts")
    plt.legend(loc='upper right', fancybox=True)
    fig_name = os.path.join(output_dir, fig_name)
    plt.savefig(fig_name)
    return counts, bins

def hist_fraction(val, idx, name='val', nbins=10, normed=False, output_dir='./'):
    counts, bins = np.histogram(val, bins=nbins)
    is_empty = np.full(len(val), False)
    for i in range(len(idx)):
        if idx[i].size == 0:
            is_empty[i] = True
    counts_detected, _ = np.histogram(val[~is_empty], bins=nbins, density=normed)
    fraction = counts_detected / counts
    fraction[np.where(np.isnan(fraction))[0]] = 0.
    plt.figure()
    plt.stairs(fraction, bins, color='r', label='completeness fraction')
    plt.xlabel(name, size='x-large')
    plt.title("Completeness Fraction")
    fig_name = os.path.join(output_dir, "completeness_fraction_%s.png"%(name))
    plt.savefig(fig_name)
    return fraction

def calculate_fraction(TU_catalog, source_catalog, output_dir, nbins=10):
    convert_catalog(TU_catalog)
    x_TU, y_TU, col_list = read_catalog(TU_catalog + '.fits', ext_num=1, ra_name="xImage", dec_name="yImage", col_list=["mag"])
    mag_TU = col_list[0]
    x_source, y_source, _ = read_catalog(source_catalog, ext_num=1, ra_name="X_IMAGE", dec_name="Y_IMAGE")
    idx1, idx2, = match_catalogs_img(x1=x_TU, y1=y_TU, x2=x_source, y2=y_source)
    counts, bins = validation_hist(val=mag_TU, idx=idx1, name="mag_injected", output_dir=output_dir)
    fraction = hist_fraction(val=mag_TU, idx=idx1, name="mag_injected", nbins=10, output_dir=output_dir)
    return counts, bins, fraction

def calculate_undetected_flux(orig_cat, mag_bins, fraction, mag_low=20.0, mag_high=26.0, image=None,  output_dir='./'):
    # Get info from original image
    hdu = fits.open(image)
    header0 = hdu[0].header
    header1 = hdu[1].header
    nx_pix, ny_pix = header0["PIXSIZE1"], header0["PIXSIZE2"]
    exp_time = header0["EXPTIME"]
    gain = header1["GAIN1"]
    chipID = int(header0["DETECTOR"][-2:])
    zp = header1["ZP"]
    hdu.close()

    # Get info from original catalog
    ra_orig, dec_orig, col_list_orig = read_catalog(orig_cat, ra_name='RA', dec_name='DEC', col_list=['Mag_Kron'])
    mag_orig = col_list_orig[0]
    nbins = len(mag_bins) - 1
    counts, _ = np.histogram(mag_orig, bins=nbins)
    
    mags = (mag_bins[:-1] + mag_bins[1:])/2.
    counts_missing = (counts / fraction) - counts
    counts_missing[np.where(np.isnan(counts_missing))[0]] = 0.
    counts_missing[np.where(np.isinf(counts_missing))[0]] = 0.
    print(counts_missing)
    print(counts_missing.sum())

    plt.figure()
    plt.stairs(counts_missing, mag_bins, color='r', label='undetected counts')
    plt.xlabel("mag_injected", size='x-large')
    plt.title("Undetected Sources")
    fig_name = os.path.join(output_dir, "undetected_sources.png")
    plt.savefig(fig_name)

    tel = Telescope()
    filter_param = FilterParam()
    filter_id, filter_type = getChipFilter(chipID=chipID)

    filt = Filter(filter_id=filter_id,
                filter_type=filter_type,
                filter_param=filter_param)
    
    undetected_flux = 0.
    for i in range(len(mags)):
        if mags[i] < mag_low or mags[i] > mag_high:
            continue
        flux_electrons = counts_missing[i] * getElectronFluxFilt(mag=mags[i], filt=filt, tel=tel)
        undetected_flux += flux_electrons

    undetected_flux /= (float(nx_pix) * float(ny_pix))
    return undetected_flux

if __name__ == "__main__":
    args = define_options().parse_args()
    counts, bins, fraction = calculate_fraction(
        TU_catalog=args.TU_catalog,
        source_catalog=args.source_catalog,
        output_dir=args.output_dir,
        nbins=20
    )
    undetected_flux = calculate_undetected_flux(
        orig_cat=args.orig_catalog,
        mag_bins=bins,
        fraction=fraction,
        image=args.image,
        output_dir=args.output_dir,
    )
    print(undetected_flux)
 No newline at end of file
+3 −3
Original line number Diff line number Diff line
@@ -41,8 +41,8 @@ def match_catalogs_sky(ra1, dec1, ra2, dec2, max_dist=0.6, others1=[], others2=[
def match_catalogs_img(x1, y1, x2, y2, max_dist=0.5, others1=[], others2=[], thresh=[]):
    cat1 = np.array([(x, y) for x,y in zip(x1, y1)])
    cat2 = np.array([(x, y) for x,y in zip(x2, y2)])
    print(np.shape(cat1))
    print(np.shape(cat2))
    # print(np.shape(cat1))
    # print(np.shape(cat2))
    tree = BallTree(cat2)
    idx1 = tree.query_radius(cat1, r = max_dist)
    tree = BallTree(cat1)
@@ -59,7 +59,7 @@ def match_catalogs_img(x1, y1, x2, y2, max_dist=0.5, others1=[], others2=[], thr

def validation_hist(val1, idx1, name="val1", nbins=10):
    counts, bins = np.histogram(val1)
    plt.stairs(bins, counts)
    plt.stairs(counts, bins)
    plt.xlabel(name, size='x-large')
    plt.savefig("detection_completeness.png")
    # plt.show()
Loading