Commit 379dedfb authored by Fang Yuedong's avatar Fang Yuedong
Browse files

adjust to new C6 catalog

parent b23efcb2
Loading
Loading
Loading
Loading
+98 −35
Original line number Original line Diff line number Diff line
@@ -5,6 +5,7 @@ import numpy as np
import h5py as h5
import h5py as h5
import healpy as hp
import healpy as hp
import astropy.constants as cons
import astropy.constants as cons
import traceback
from astropy.coordinates import spherical_to_cartesian
from astropy.coordinates import spherical_to_cartesian
from astropy.table import Table
from astropy.table import Table
from scipy import interpolate
from scipy import interpolate
@@ -14,6 +15,11 @@ from ObservationSim.MockObject import CatalogBase, Star, Galaxy, Quasar
from ObservationSim.MockObject._util import tag_sed, getObservedSED, getABMAG, integrate_sed_bandpass, comoving_dist
from ObservationSim.MockObject._util import tag_sed, getObservedSED, getABMAG, integrate_sed_bandpass, comoving_dist
from ObservationSim.Astrometry.Astrometry_util import on_orbit_obs_position
from ObservationSim.Astrometry.Astrometry_util import on_orbit_obs_position


# (TEST)
from astropy.cosmology import FlatLambdaCDM
from astropy import constants
from astropy import units as U

try:
try:
    import importlib.resources as pkg_resources
    import importlib.resources as pkg_resources
except ImportError:
except ImportError:
@@ -22,26 +28,40 @@ except ImportError:


NSIDE = 128
NSIDE = 128


def get_bundleIndex(healpixID_ring, bundleOrder=4, healpixOrder=7):
    assert NSIDE == 2**healpixOrder
    shift = healpixOrder - bundleOrder
    shift = 2*shift

    nside_bundle = 2**bundleOrder
    nside_healpix= 2**healpixOrder

    healpixID_nest= hp.ring2nest(nside_healpix, healpixID_ring)
    bundleID_nest = (healpixID_nest >> shift)
    bundleID_ring = hp.nest2ring(nside_bundle, bundleID_nest)

    return bundleID_ring

class C6_Catalog(CatalogBase):
class C6_Catalog(CatalogBase):
    def __init__(self, config, chip, pointing, chip_output, filt, **kwargs):
    def __init__(self, config, chip, pointing, chip_output, filt, **kwargs):
        super().__init__()
        super().__init__()
        self.cat_dir = os.path.join(config["data_dir"], config["input_path"]["cat_dir"])
        self.cat_dir = os.path.join(config["data_dir"], config["input_path"]["cat_dir"])
        self.seed_Av = config["random_seeds"]["seed_Av"]
        self.seed_Av = config["random_seeds"]["seed_Av"]


        # (TEST)
        self.cosmo = FlatLambdaCDM(H0=67.66, Om0=0.3111)

        self.chip_output = chip_output
        self.chip_output = chip_output
        self.filt = filt
        self.filt = filt
        self.logger = chip_output.logger
        self.logger = chip_output.logger


        with pkg_resources.path('Catalog.data', 'SLOAN_SDSS.g.fits') as filter_path:
        with pkg_resources.path('Catalog.data', 'SLOAN_SDSS.g.fits') as filter_path:
            self.normF_star = Table.read(str(filter_path))
            self.normF_star = Table.read(str(filter_path))
        # with pkg_resources.path('Catalog.data', 'lsst_throuput_g.fits') as filter_path:
        #     self.normF_galaxy = Table.read(str(filter_path))
        
        
        self.config = config
        self.config = config
        self.chip = chip
        self.chip = chip
        self.pointing = pointing
        self.pointing = pointing


        # (DEBUG)
        self.max_size = 0.
        self.max_size = 0.


        if "star_cat" in config["input_path"] and config["input_path"]["star_cat"] and not config["run_option"]["galaxy_only"]:
        if "star_cat" in config["input_path"] and config["input_path"]["star_cat"] and not config["run_option"]["galaxy_only"]:
@@ -80,8 +100,13 @@ class C6_Catalog(CatalogBase):
        ra_min, ra_max, dec_min, dec_max = self.sky_coverage.xmin, self.sky_coverage.xmax, self.sky_coverage.ymin, self.sky_coverage.ymax
        ra_min, ra_max, dec_min, dec_max = self.sky_coverage.xmin, self.sky_coverage.xmax, self.sky_coverage.ymin, self.sky_coverage.ymax
        ra = np.deg2rad(np.array([ra_min, ra_max, ra_max, ra_min]))
        ra = np.deg2rad(np.array([ra_min, ra_max, ra_max, ra_min]))
        dec = np.deg2rad(np.array([dec_max, dec_max, dec_min, dec_min]))
        dec = np.deg2rad(np.array([dec_max, dec_max, dec_min, dec_min]))
        vertices = spherical_to_cartesian(1., dec, ra)
        # vertices = spherical_to_cartesian(1., dec, ra)
        self.pix_list = hp.query_polygon(NSIDE, np.array(vertices).T, inclusive=True)
        self.pix_list = hp.query_polygon(
            NSIDE,
            hp.ang2vec(np.radians(90.) - dec, ra),
            inclusive=True
        )
        # self.pix_list = hp.query_polygon(NSIDE, np.array(vertices).T, inclusive=True)
        if self.logger is not None:
        if self.logger is not None:
            msg = str(("HEALPix List: ", self.pix_list))
            msg = str(("HEALPix List: ", self.pix_list))
            self.logger.info(msg)
            self.logger.info(msg)
@@ -144,16 +169,21 @@ class C6_Catalog(CatalogBase):
            )
            )


        for igals in range(ngals):
        for igals in range(ngals):
            # (TEST)
            # # (TEST)
            # if igals > 100:
            # if igals > 100:
            #     break
            #     break
            # if igals != 1186:
            #     continue
            
            
            param = self.initialize_param()
            param = self.initialize_param()
            param['ra'] = ra_arr[igals]
            param['ra'] = ra_arr[igals]
            param['dec'] = dec_arr[igals]
            param['dec'] = dec_arr[igals]
            param['ra_orig'] = gals['ra'][igals]
            param['ra_orig'] = gals['ra'][igals]
            param['dec_orig'] = gals['dec'][igals]
            param['dec_orig'] = gals['dec'][igals]
            # param['mag_use_normal'] = gals['mag_true_g_lsst'][igals]
            param['mag_use_normal'] = gals['mag_csst_%s'%(self.filt.filter_type)][igals]
            if self.filt.is_too_dim(mag=param['mag_use_normal'], margin=self.config["obs_setting"]["mag_lim_margin"]):
                continue

            # if param['mag_use_normal'] >= 26.5:
            # if param['mag_use_normal'] >= 26.5:
            #     continue
            #     continue
            param['z'] = gals['redshift'][igals]
            param['z'] = gals['redshift'][igals]
@@ -165,11 +195,16 @@ class C6_Catalog(CatalogBase):
            param['e2'] = gals['ellipticity_true'][igals][1]
            param['e2'] = gals['ellipticity_true'][igals][1]
            
            
            # For shape calculation
            # For shape calculation
            
            param['ell_total'] = np.sqrt(param['e1']**2 + param['e2']**2)
            if param['ell_total'] > 0.9:
                continue
            param['e1_disk'] = param['e1']
            param['e1_disk'] = param['e1']
            param['e2_disk'] = param['e2']
            param['e2_disk'] = param['e2']
            param['e1_bulge'] = param['e1']
            param['e1_bulge'] = param['e1']
            param['e2_bulge'] = param['e2']
            param['e2_bulge'] = param['e2']



            param['delta_ra'] = 0
            param['delta_ra'] = 0
            param['delta_dec'] = 0
            param['delta_dec'] = 0


@@ -211,7 +246,7 @@ class C6_Catalog(CatalogBase):
            # TEMP
            # TEMP
            self.ids += 1
            self.ids += 1
            # param['id'] = self.ids
            # param['id'] = self.ids
            param['id'] = str(pix_id) + '%02d'%(cat_id) + '%08d'%(igals)
            param['id'] = '%06d'%(int(pix_id)) + '%06d'%(cat_id) + '%08d'%(igals)
            
            
            if param['star'] == 0:
            if param['star'] == 0:
                obj = Galaxy(param, self.rotation, logger=self.logger)
                obj = Galaxy(param, self.rotation, logger=self.logger)
@@ -264,8 +299,8 @@ class C6_Catalog(CatalogBase):
                input_time_str=time_str
                input_time_str=time_str
            )
            )
        for istars in range(nstars):
        for istars in range(nstars):
            # (TEST)
            # # (TEST)
            # if istars > 10:
            # if istars > 100:
            #     break
            #     break


            param = self.initialize_param()
            param = self.initialize_param()
@@ -283,7 +318,6 @@ class C6_Catalog(CatalogBase):
            # if param['mag_use_normal'] >= 26.5:
            # if param['mag_use_normal'] >= 26.5:
            #     continue
            #     continue
            self.ids += 1
            self.ids += 1
            # param['id'] = self.ids
            param['id'] = stars['sourceID'][istars]
            param['id'] = stars['sourceID'][istars]
            param['sed_type'] = stars['sourceID'][istars]
            param['sed_type'] = stars['sourceID'][istars]
            param['model_tag'] = stars['model_tag'][istars]
            param['model_tag'] = stars['model_tag'][istars]
@@ -316,15 +350,28 @@ class C6_Catalog(CatalogBase):
                    print(e)
                    print(e)
        
        
        if "galaxy_cat" in self.config["input_path"] and self.config["input_path"]["galaxy_cat"] and not self.config["run_option"]["star_only"]:
        if "galaxy_cat" in self.config["input_path"] and self.config["input_path"]["galaxy_cat"] and not self.config["run_option"]["star_only"]:
            for i in range(76):
            # for i in range(76):
                file_path = os.path.join(self.galaxy_path, "galaxies%04d_C6.h5"%(i))
            #     file_path = os.path.join(self.galaxy_path, "galaxies%04d_C6.h5"%(i))
                gals_cat = h5.File(file_path, 'r')['galaxies']
            #     gals_cat = h5.File(file_path, 'r')['galaxies']
            #     for pix in self.pix_list:
            #         try:
            #             gals = gals_cat[str(pix)]
            #             self._load_gals(gals, pix_id=pix, cat_id=i)
            #             del gals
            #         except Exception as e:
            #             traceback.print_exc()
            #             self.logger.error(str(e))
            #             print(e)
            for pix in self.pix_list:
            for pix in self.pix_list:
                try:
                try:
                    bundleID  = get_bundleIndex(pix)
                    file_path = os.path.join(self.galaxy_path, "galaxies_C6_bundle{:06}.h5".format(bundleID))
                    gals_cat = h5.File(file_path, 'r')['galaxies']
                    gals = gals_cat[str(pix)]
                    gals = gals_cat[str(pix)]
                        self._load_gals(gals, pix_id=pix, cat_id=i)
                    self._load_gals(gals, pix_id=pix, cat_id=bundleID)
                    del gals
                    del gals
                except Exception as e:
                except Exception as e:
                    traceback.print_exc()
                    self.logger.error(str(e))
                    self.logger.error(str(e))
                    print(e)
                    print(e)
        if self.logger is not None:
        if self.logger is not None:
@@ -333,6 +380,19 @@ class C6_Catalog(CatalogBase):
        else:
        else:
            print("number of objects in catalog: ", len(self.objs))
            print("number of objects in catalog: ", len(self.objs))


    # (TEST)
    # def convert_mags(self, obj):
    #     spec = np.matmul(obj.coeff, self.pcs.T)
    #     lamb = self.lamb_gal * U.angstrom
    #     unit_sed = ((lamb * lamb)/constants.c*U.erg/U.second/U.cm**2/U.angstrom).to(U.jansky)
    #     unit_sed = unit_sed.value
    #     lamb = lamb.value
    #     lamb *= (1 + obj.z)
    #     spec *= (1 + obj.z)
    #     mags = -2.5 * np.log10(unit_sed*spec) + 8.9
    #     mags += self.cosmo.distmod(obj.z).value
    #     return lamb, mags



    def load_sed(self, obj, **kwargs):
    def load_sed(self, obj, **kwargs):
        if obj.type == 'star':
        if obj.type == 'star':
@@ -344,11 +404,13 @@ class C6_Catalog(CatalogBase):
                feh=obj.param['feh']
                feh=obj.param['feh']
            )
            )
        elif obj.type == 'galaxy' or obj.type == 'quasar':
        elif obj.type == 'galaxy' or obj.type == 'quasar':
            dist_L_pc = (1 + obj.z) * comoving_dist(z=obj.z)[0]
            # dist_L_pc = (1 + obj.z) * comoving_dist(z=obj.z)[0]
            factor = (10 / dist_L_pc)**2
            # factor = (10 / dist_L_pc)**2
            factor = 10**(-.4 * self.cosmo.distmod(obj.z).value)
            flux = np.matmul(self.pcs, obj.coeff) * factor
            flux = np.matmul(self.pcs, obj.coeff) * factor
            if np.any(flux < 0):
            # if np.any(flux < 0):
                raise ValueError("Glaxy %s: negative SED fluxes"%obj.id)
            #     raise ValueError("Glaxy %s: negative SED fluxes"%obj.id)
            flux[flux < 0] = 0.
            sedcat = np.vstack((self.lamb_gal, flux)).T
            sedcat = np.vstack((self.lamb_gal, flux)).T
            sed_data = getObservedSED(
            sed_data = getObservedSED(
                sedCat=sedcat,
                sedCat=sedcat,
@@ -368,15 +430,16 @@ class C6_Catalog(CatalogBase):
        all_sed = y * lamb / (cons.h.value * cons.c.value) * 1e-13
        all_sed = y * lamb / (cons.h.value * cons.c.value) * 1e-13
        sed = Table(np.array([lamb, all_sed]).T, names=('WAVELENGTH', 'FLUX'))
        sed = Table(np.array([lamb, all_sed]).T, names=('WAVELENGTH', 'FLUX'))
        
        
        if obj.type == 'galaxy' or obj.type == 'quasar':
        # if obj.type == 'galaxy' or obj.type == 'quasar':
            # integrate to get the magnitudes
        #     # integrate to get the magnitudes
            sed_photon = np.array([sed['WAVELENGTH'], sed['FLUX']]).T
        #     sed_photon = np.array([sed['WAVELENGTH'], sed['FLUX']]).T
            # print("sed_photon = ", sed_photon)
        #     sed_photon = galsim.LookupTable(x=np.array(sed_photon[:, 0]), f=np.array(sed_photon[:, 1]), interpolant='nearest')
            sed_photon = galsim.LookupTable(x=np.array(sed_photon[:, 0]), f=np.array(sed_photon[:, 1]), interpolant='nearest')
        #     sed_photon = galsim.SED(sed_photon, wave_type='A', flux_type='1', fast=False)
            sed_photon = galsim.SED(sed_photon, wave_type='A', flux_type='1', fast=False)
        #     interFlux = integrate_sed_bandpass(sed=sed_photon, bandpass=self.filt.bandpass_full)
            interFlux = integrate_sed_bandpass(sed=sed_photon, bandpass=self.filt.bandpass_full)
        #     obj.param['mag_use_normal'] = getABMAG(interFlux, self.filt.bandpass_full)
            obj.param['mag_use_normal'] = getABMAG(interFlux, self.filt.bandpass_full)
        #     # print("mag_use_normal = %.3f"%obj.param['mag_use_normal'])
            # print("mag_use_normal = %.3f"%obj.param['mag_use_normal'])
        #     # mag = getABMAG(interFlux, self.filt.bandpass_full)
        #     # print("mag diff = %.3f"%(mag - obj.param['mag_use_normal']))
        
        
        del wave
        del wave
        del flux
        del flux
+1 −0
Original line number Original line Diff line number Diff line
@@ -229,6 +229,7 @@ class Chip(FocalPlane):




    def getSkyCoverage(self, wcs):
    def getSkyCoverage(self, wcs):
        # print("In getSkyCoverage: xmin = %.3f, xmax = %.3f, ymin = %.3f, ymax = %.3f"%(self.bound.xmin, self.bound.xmax, self.bound.ymin, self.bound.ymax))
        return super().getSkyCoverage(wcs, self.bound.xmin, self.bound.xmax, self.bound.ymin, self.bound.ymax)
        return super().getSkyCoverage(wcs, self.bound.xmin, self.bound.xmax, self.bound.ymin, self.bound.ymax)




+33 −26
Original line number Original line Diff line number Diff line
@@ -137,6 +137,7 @@ class Galaxy(MockObject):
                # return False
                # return False
                continue
                continue
            nphotons_sum += nphotons
            nphotons_sum += nphotons

            # print("nphotons_sub-band_%d = %.2f"%(i, nphotons))
            # print("nphotons_sub-band_%d = %.2f"%(i, nphotons))


            psf, pos_shear = psf_model.get_PSF(chip=chip, pos_img=pos_img, bandpass=bandpass, folding_threshold=folding_threshold)
            psf, pos_shear = psf_model.get_PSF(chip=chip, pos_img=pos_img, bandpass=bandpass, folding_threshold=folding_threshold)
@@ -171,8 +172,10 @@ class Galaxy(MockObject):
            # photons_list.append(photons)
            # photons_list.append(photons)


            stamp = gal.drawImage(wcs=real_wcs_local, method='phot', offset=offset, save_photons=True)
            stamp = gal.drawImage(wcs=real_wcs_local, method='phot', offset=offset, save_photons=True)
            
            xmax = max(xmax, stamp.xmax)
            xmax = max(xmax, stamp.xmax)
            ymax = max(ymax, stamp.ymax)
            ymax = max(ymax, stamp.ymax)

            photons = stamp.photons
            photons = stamp.photons
            photons.x += x_nominal
            photons.x += x_nominal
            photons.y += y_nominal
            photons.y += y_nominal
@@ -184,6 +187,8 @@ class Galaxy(MockObject):
        stamp.wcs = real_wcs_local
        stamp.wcs = real_wcs_local
        stamp.setCenter(x_nominal, y_nominal)
        stamp.setCenter(x_nominal, y_nominal)
        bounds = stamp.bounds & galsim.BoundsI(0, chip.npix_x - 1, 0, chip.npix_y - 1)
        bounds = stamp.bounds & galsim.BoundsI(0, chip.npix_x - 1, 0, chip.npix_y - 1)

        if bounds.area() > 0:
            chip.img.setOrigin(0, 0)
            chip.img.setOrigin(0, 0)
            stamp[bounds] = chip.img[bounds]
            stamp[bounds] = chip.img[bounds]


@@ -237,7 +242,7 @@ class Galaxy(MockObject):


    def drawObj_slitless(self, tel, pos_img, psf_model, bandpass_list, filt, chip, nphotons_tot=None, g1=0, g2=0,
    def drawObj_slitless(self, tel, pos_img, psf_model, bandpass_list, filt, chip, nphotons_tot=None, g1=0, g2=0,
                         exptime=150., normFilter=None, grating_split_pos=3685):
                         exptime=150., normFilter=None, grating_split_pos=3685):

        if normFilter is not None:
            norm_thr_rang_ids = normFilter['SENSITIVITY'] > 0.001
            norm_thr_rang_ids = normFilter['SENSITIVITY'] > 0.001
            sedNormFactor = getNormFactorForSpecWithABMAG(ABMag=self.param['mag_use_normal'], spectrum=self.sed,
            sedNormFactor = getNormFactorForSpecWithABMAG(ABMag=self.param['mag_use_normal'], spectrum=self.sed,
                                                        norm_thr=normFilter,
                                                        norm_thr=normFilter,
@@ -245,6 +250,8 @@ class Galaxy(MockObject):
                                                        eWave=np.ceil(normFilter[norm_thr_rang_ids][-1][0]))
                                                        eWave=np.ceil(normFilter[norm_thr_rang_ids][-1][0]))
            if sedNormFactor == 0:
            if sedNormFactor == 0:
                return False
                return False
        else:
            sedNormFactor = 1.
        normalSED = Table(np.array([self.sed['WAVELENGTH'], self.sed['FLUX'] * sedNormFactor]).T,
        normalSED = Table(np.array([self.sed['WAVELENGTH'], self.sed['FLUX'] * sedNormFactor]).T,
                          names=('WAVELENGTH', 'FLUX'))
                          names=('WAVELENGTH', 'FLUX'))


+15 −11
Original line number Original line Diff line number Diff line
import galsim
import galsim
import numpy as np
import numpy as np
import astropy.constants as cons
import astropy.constants as cons
from astropy import wcs
from astropy.table import Table
from astropy.table import Table


from ObservationSim.MockObject._util import magToFlux, VC_A, convolveGaussXorders
from ObservationSim.MockObject._util import magToFlux, VC_A, convolveGaussXorders
@@ -99,8 +100,6 @@ class MockObject(object):
        dy = y - self.y_nominal
        dy = y - self.y_nominal
        self.offset = galsim.PositionD(dx, dy)
        self.offset = galsim.PositionD(dx, dy)


        from astropy import wcs

        if img_header is not None:
        if img_header is not None:
            self.real_wcs = galsim.FitsWCS(header=img_header)
            self.real_wcs = galsim.FitsWCS(header=img_header)
        else:
        else:
@@ -224,6 +223,11 @@ class MockObject(object):
        stamp.wcs = real_wcs_local
        stamp.wcs = real_wcs_local
        stamp.setCenter(x_nominal, y_nominal)
        stamp.setCenter(x_nominal, y_nominal)
        bounds = stamp.bounds & galsim.BoundsI(0, chip.npix_x - 1, 0, chip.npix_y - 1)
        bounds = stamp.bounds & galsim.BoundsI(0, chip.npix_x - 1, 0, chip.npix_y - 1)
        
        # # (DEBUG)
        # print("stamp bounds: ", stamp.bounds)
        # print(bounds)
        if bounds.area() > 0:
            chip.img.setOrigin(0, 0)
            chip.img.setOrigin(0, 0)
            stamp[bounds] = chip.img[bounds]
            stamp[bounds] = chip.img[bounds]
            for i in range(len(photons_list)):
            for i in range(len(photons_list)):
+4 −1
Original line number Original line Diff line number Diff line
@@ -455,7 +455,10 @@ def getObservedSED(sedCat, redshift=0.0, av=0.0, redden=0):
    sw, sf = sedCat[:,0], sedCat[:,1]
    sw, sf = sedCat[:,0], sedCat[:,1]
    # reddening
    # reddening
    sf = reddening(sw, sf, av=av, model=redden)
    sf = reddening(sw, sf, av=av, model=redden)
    sw, sf = sw*z, sf*(z**3)
    # sw, sf = sw*z, sf*(z**3)
    sw, sf = sw*z, sf/z
    # sw, sf = sw*z, sf

    # lyman forest correction
    # lyman forest correction
    sf = lyman_forest(sw, sf, redshift)
    sf = lyman_forest(sw, sf, redshift)
    isedObs = (sw.copy(), sf.copy())
    isedObs = (sw.copy(), sf.copy())
Loading