Repository navigation
add scripts testing full scene drawing residuals+timing #190
Description
Activity
Here is my first draft of how to do the scene drawing , I tried to follow the procedure in descwl-shear-sims. @beckermr was that your suggestion?
Here is the galsim code, which appears to work without issues:
def get_convolved_object_galsim( flux_d, flux_b, hlr_b, hlr_d, q_b, q_d, beta, *, psf_hlr=0.7, ) -> galsim.GSObject: components = [] # disk disk = galsim.Exponential(flux=flux_d, half_light_radius=hlr_d).shear( q=q_d, beta=beta * galsim.degrees ) components.append(disk) # bulge bulge = galsim.Spergel(nu=-0.6, flux=flux_b, half_light_radius=hlr_b).shear( q=q_b, beta=beta * galsim.degrees ) components.append(bulge) galaxy = galsim.Add(components) # psf # psf = xgalsim.Moffat(2, flux=1.0, scale_radius=0.7) # psf = xgalsim.Moffat(2, flux=1.0, half_light_radius=0.7) psf = galsim.Gaussian(flux=1.0, half_light_radius=psf_hlr) gal_conv = galsim.Convolve([galaxy, psf]) return gal_conv def draw_galsim(galaxy_params: dict, *, slen: int): # create big image image = galsim.Image(ncol=slen, nrow=slen, scale=0.2, dtype=np.float64) wcs = image.wcs n_sources = len(galaxy_params["flux_d"]) for n in range(n_sources): _gal_params = {k: v[n].item() for k, v in galaxy_params.items()} x = _gal_params.pop("x") y = _gal_params.pop("y") image_pos = galsim.PositionD(x=x, y=y) local_wcs = wcs.local(image_pos=image_pos) convolved_object = get_convolved_object_galsim(**_gal_params) stamp = convolved_object.drawImage( center=image_pos, wcs=local_wcs, dtype=image.dtype ) b = stamp.bounds & image.bounds if b.isDefined(): image[b] += stamp[b] # add noise background = get_default_lsst_background() x = image.array + background x = np.random.normal(loc=x, scale=np.sqrt(background), size=x.shape) x -= background # background subtracted return x
where
galaxy_paramsis just a dictionary with all the galaxy properties for all galaxies in the full scene.Here is my attempt at the jax galsim version of the code:
def draw_stamp_xgalsim( galaxy_params: dict, image_pos: xgalsim.PositionD, local_wcs: xgalsim.PixelScale, *, fft_size: int, slen: int, ) -> jax.Array: gsparams = xgalsim.GSParams(minimum_fft_size=fft_size, maximum_fft_size=fft_size) convolved_object = get_bd_xgalsim(**galaxy_params).withGSParams(gsparams) stamp = convolved_object.drawImage( nx=slen, ny=slen, center=image_pos, wcs=local_wcs, dtype=jnp.float64 ) return stamp def draw_xgalsim(galaxy_params: dict, *, slen: int, draw_vmap_func: Callable): # create big image image = xgalsim.Image(ncol=slen, nrow=slen, scale=0.2, dtype=np.float64) wcs = image.wcs _position_fnc = lambda x, y: xgalsim.PositionD(x=x, y=y) _local = lambda x: wcs.local(image_pos=x) xs = galaxy_params.pop("x") ys = galaxy_params.pop("y") image_positions = vmap(_position_fnc)(xs, ys) local_wcss = vmap(_local)(image_positions) stamps = draw_vmap_func(galaxy_params, image_positions, local_wcss) # TODO: put together the stamps and add to image (in the CPU?) # add noise background = get_default_lsst_background() x = image.array + background x = np.random.normal(loc=x, scale=np.sqrt(background), size=x.shape) x -= background # background subtracted return x
where for example
draw_vmap_func = jax.jit(jax.vmap(partial(draw_stamp_xgalsim, slen=52, fft_size=256)))
The jax galsim code crashes when drawing all the stamps vectorized. It seems there is a problem with the checks that are done to the bounds that fail because of the abstract tracer:
File ~/code/JAX-GalSim/jax_galsim/image.py:490, in Image.subImage(self, bounds) 486 if not self.bounds.isDefined(): 487 raise _galsim.GalSimUndefinedBoundsError( 488 "Attempt to access subImage of undefined image" 489 ) --> [490](https://file+.vscode-resource.vscode-cdn.net/Users/imendoza/code/JAX-GalSim/notebooks/~/code/JAX-GalSim/jax_galsim/image.py:490) if not self.bounds.includes(bounds): 491 raise _galsim.GalSimBoundsError( 492 "Attempt to access subImage not (fully) in image", bounds, self.bounds 493 ) 494 i1 = bounds.ymin - self.ymin File ~/code/JAX-GalSim/jax_galsim/bounds.py:112, in Bounds.includes(self, *args) 109 if isinstance(args[0], Bounds): 110 b = args[0] 111 return ( --> [112](https://file+.vscode-resource.vscode-cdn.net/Users/imendoza/code/JAX-GalSim/notebooks/~/code/JAX-GalSim/jax_galsim/bounds.py:112) self.isDefined() 113 and b.isDefined() 114 and self.xmin <= b.xmin 115 and self.xmax >= b.xmax 116 and self.ymin <= b.ymin 117 and self.ymax >= b.ymax 118 ) 119 elif isinstance(args[0], Position): 120 p = args[0] [... skipping hidden 1 frame] File ~/code/JAX-GalSim/.venv/lib/python3.13/site-packages/jax/_src/core.py:[1829](https://file+.vscode-resource.vscode-cdn.net/Users/imendoza/code/JAX-GalSim/notebooks/~/code/JAX-GalSim/.venv/lib/python3.13/site-packages/jax/_src/core.py:1829), in concretization_function_error.<locals>.error(self, arg) 1828 def error(self, arg): -> 1829 raise TracerBoolConversionError(arg)Do you think I different approach is warranted or there is away to get around this bound checks? Thanks!
I think it is this call:
b = stamp.bounds & image.bounds if b.isDefined(): image[b] += stamp[b]I think we can either convert to a
condoperator or you can add all of the objects first and then render.Reacted by Ismael MendozaYou should also be sure to declare the fft_size parameter as a static arg for the JIT decorator
Hi, I have cooked this Pay attention that
draw_stamp_xgalsimhas a first argument the result of a convolution PSF,galaxy.def get_convolved_object_xgalsim( flux_d, flux_b, hlr_b, hlr_d, q_b, q_d, beta, *, psf_hlr=0.7, ) -> xgalsim.GSObject: components = [] # disk disk = xgalsim.Exponential(flux=flux_d, half_light_radius=hlr_d).shear( q=q_d, beta=beta * xgalsim.degrees ) components.append(disk) # bulge bulge = xgalsim.Spergel(nu=-0.6, flux=flux_b, half_light_radius=hlr_b).shear( q=q_b, beta=beta * xgalsim.degrees ) components.append(bulge) galaxy = xgalsim.Add(components) # psf psf = xgalsim.Gaussian(flux=1.0, half_light_radius=psf_hlr) gal_conv = xgalsim.Convolve([galaxy, psf]) return gal_conv def draw_stamp_xgalsim( convolved_object: xgalsim.Convolution, image_pos: xgalsim.PositionD, local_wcs: xgalsim.PixelScale, *, fft_size: int, slen: int, ) -> jax.Array: gsparams = xgalsim.GSParams(minimum_fft_size=fft_size, maximum_fft_size=fft_size) jax.debug.print("image_pos: {}, local_wcs: {}",image_pos,local_wcs) stamp = convolved_object.drawImage( nx=slen, ny=slen, center=image_pos, wcs=local_wcs, dtype=jnp.float64 ) return stamp
and use it first w/o vmap/jit
#create an object obj= get_convolved_object_xgalsim(flux_d=10**3, flux_b=100, hlr_b=10, hlr_d=20, q_b=0.5, q_d=0.9, beta=jnp.pi/4) # create a big image image = xgalsim.Image(ncol=512, nrow=512, scale=0.2, dtype=jnp.float64) wcs = image.wcs # some parameters fft_size=256 slen= 52 # the positions of three instances xs=[256,100,300] ys=[256,200,400] # for i in range(len(xs)): image_pos = xgalsim.PositionD(x=xs[i], y=ys[i]) local_wcs = wcs.local(image_pos=image_pos) stamp = draw_stamp_xgalsim(obj, image_pos, local_wcs, fft_size=fft_size, slen=slen) b = stamp.bounds & image.bounds if b.isDefined(): image[b] += stamp[b]
This gives as output
image_pos: galsim.PositionD(256.0,256.0), local_wcs: galsim.PixelScale(0.2) image_pos: galsim.PositionD(100.0,200.0), local_wcs: galsim.PixelScale(0.2) image_pos: galsim.PositionD(300.0,400.0), local_wcs: galsim.PixelScale(0.2)and
imshow(image.array,cmap="jet")
Now with the "obj" and the xs,ys positions trying to make a similar code as @ismael-mendoza
image_positions = jax.vmap(lambda x, y: xgalsim.PositionD(x,y))(xs, ys)
gives
galsim.PositionD(x=(256.0, 100.0, 300.0), y=(256.0, 200.0, 400.0))which is the normal way to think of JAX (struct-of-arrays), we can convert it into a list-of-PositionsD but this is not the point.
thenlocal_wcss = jax.vmap(lambda x: wcs.local(image_pos=x))(image_positions)
(I cannot print anything here there is a crash 'not all arguments converted during string formatting")
Now comes the jit/vmap try: first fix the slen, fft_size pramaters as static
jit_draw_stamp_xgalsim = jax.jit(draw_stamp_xgalsim, static_argnames=("slen","fft_size"))
thenn vmap a partial version fixing the single object to be used here and the values of slen, fft_size pramaters
draw_vmap_func = jax.vmap(partial(jit_draw_stamp_xgalsim,convolved_object=obj,slen=52,fft_size=256))
Then we would expect the following to draxw the three stamps
draw_vmap_func(image_pos=image_positions, local_wcs=local_wcss)But I have a crash certainly related to the fft_size even if I make it static
File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:783](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/gsobject.py#line=782), in GSObject.drawImage(self, image, nx, ny, bounds, scale, wcs, dtype, method, area, exptime, gain, add_to_image, center, use_true_center, offset, n_photons, rng, max_extra_noise, poisson_flux, sensor, photon_ops, n_subsample, maxN, save_photons, bandpass, setup_only, surface_ops) 781 added_photons = prof.drawReal(image, add_to_image) 782 else: --> 783 added_photons = prof.drawFFT(image, add_to_image) 785 image.added_flux = added_photons [/](https://jupyterhub.idris.fr/) flux_scale 786 # Restore the original center and wcs File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:928](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/gsobject.py#line=927), in GSObject.drawFFT(self, image, add_to_image) 923 if image.wcs is None or not image.wcs.isPixelScale(): 924 raise _galsim.GalSimValueError( 925 "drawFFT requires an image with a PixelScale wcs", image 926 ) --> 928 kimage, wrap_size = self.drawFFT_makeKImage(image) 929 kimage = self._drawKImage(kimage) 930 return self.drawFFT_finish(image, kimage, wrap_size, add_to_image) File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:864](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/gsobject.py#line=863), in GSObject.drawFFT_makeKImage(self, image) 861 N = jnp.max(jnp.array([N, image_N])) 863 # Round up to a good size for making FFTs: --> 864 N = image.good_fft_size(N) 866 # Make sure we hit the minimum size specified in the gsparams. 867 N = max(N, self.gsparams.minimum_fft_size) File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/image.py:786](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/image.py#line=785), in Image.good_fft_size(cls, input_size) 782 # Reference from GalSim C++ 783 # https://github.com/GalSim-developers/GalSim/blob/ece3bd32c1ae6ed771f2b489c5ab1b25729e0ea4/src/Image.cpp#L1009 784 # Reduce slightly to eliminate potential rounding errors: 785 insize = (1.0 - 1.0e-5) * input_size --> 786 log2n = math.log(2.0) * math.ceil(math.log(insize) [/](https://jupyterhub.idris.fr/) math.log(2.0)) 787 log2n3 = math.log(3.0) + math.log(2.0) * math.ceil( 788 (math.log(insize) - math.log(3.0)) [/](https://jupyterhub.idris.fr/) math.log(2.0) 789 ) 790 log2n3 = max(log2n3, math.log(6.0)) # must be even number [... skipping hidden 1 frame] File [/lustre/fswork/projects/rech/ixh/ufd72rp/.local_jax_galsim/lib/python3.12/site-packages/jax/_src/core.py:1835](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/.local_jax_galsim/lib/python3.12/site-packages/jax/_src/core.py#line=1834), in concretization_function_error.<locals>.error(self, arg) 1834 def error(self, arg): -> 1835 raise ConcretizationTypeError(arg, fname_context) ConcretizationTypeError: Abstract tracer value encountered where concrete value is expected: traced array with shape float64[] The problem arose with the `float` function. If trying to convert the data type of a value, try using `x.astype(float)` or `jnp.array(x, float)` instead. The error occurred while tracing the function draw_stamp_xgalsim at [/tmp/ipykernel_760504/2509194331.py:35](https://jupyterhub.idris.fr/tmp/ipykernel_760504/2509194331.py#line=34) for jit. This concrete value was not available in Python because it depends on the values of the arguments convolved_object[0]['obj_list'][0][0]['obj_list'][0][0][0]['scale_radius'], convolved_object[0]['obj_list'][0][0]['obj_list'][0][1]['jac'], convolved_object[0]['obj_list'][0][0]['obj_list'][0][1]['offset'][0], convolved_object[0]['obj_list'][0][0]['obj_list'][0][1]['offset'][1], convolved_object[0]['obj_list'][0][0]['obj_list'][1][0][0]['nu'], convolved_object[0]['obj_list'][0][0]['obj_list'][1][0][0]['scale_radius'], convolved_object[0]['obj_list'][0][0]['obj_list'][1][1]['jac'], convolved_object[0]['obj_list'][0][0]['obj_list'][1][1]['offset'][0], convolved_object[0]['obj_list'][0][0]['obj_list'][1][1]['offset'][1], convolved_object[0]['obj_list'][1][0]['sigma'], image_pos[0], image_pos[1], and local_wcs[0]['scale'].In fact w/o the vmaping
jax.jit(draw_stamp_xgalsim, static_argnames=('slen','fft_size'))(convolved_object=obj, image_pos=xgalsim.PositionD(126,126), local_wcs=xgalsim.PixelScale(0.2), fft_size=256, slen=52, )
produces an error
-------------------------------------------------------------------------- ConcretizationTypeError Traceback (most recent call last) Cell In[191], line 1 ----> 1 jax.jit(draw_stamp_xgalsim, static_argnames=('slen','fft_size'))(convolved_object=obj, 2 image_pos=xgalsim.PositionD(126,126), 3 local_wcs=xgalsim.PixelScale(0.2), 4 fft_size=256, 5 slen=52, 6 ) [... skipping hidden 13 frame] Cell In[139], line 46, in draw_stamp_xgalsim(convolved_object, image_pos, local_wcs, fft_size, slen) 43 gsparams = xgalsim.GSParams(minimum_fft_size=fft_size, maximum_fft_size=fft_size) 45 jax.debug.print("image_pos: {}, local_wcs: {}",image_pos,local_wcs) ---> 46 stamp = convolved_object.drawImage( 47 nx=slen, ny=slen, center=image_pos, wcs=local_wcs, dtype=jnp.float64 48 ) 49 return stamp File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:783](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/gsobject.py#line=782), in GSObject.drawImage(self, image, nx, ny, bounds, scale, wcs, dtype, method, area, exptime, gain, add_to_image, center, use_true_center, offset, n_photons, rng, max_extra_noise, poisson_flux, sensor, photon_ops, n_subsample, maxN, save_photons, bandpass, setup_only, surface_ops) 781 added_photons = prof.drawReal(image, add_to_image) 782 else: --> 783 added_photons = prof.drawFFT(image, add_to_image) 785 image.added_flux = added_photons [/](https://jupyterhub.idris.fr/) flux_scale 786 # Restore the original center and wcs File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:928](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/gsobject.py#line=927), in GSObject.drawFFT(self, image, add_to_image) 923 if image.wcs is None or not image.wcs.isPixelScale(): 924 raise _galsim.GalSimValueError( 925 "drawFFT requires an image with a PixelScale wcs", image 926 ) --> 928 kimage, wrap_size = self.drawFFT_makeKImage(image) 929 kimage = self._drawKImage(kimage) 930 return self.drawFFT_finish(image, kimage, wrap_size, add_to_image) File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:864](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/gsobject.py#line=863), in GSObject.drawFFT_makeKImage(self, image) 861 N = jnp.max(jnp.array([N, image_N])) 863 # Round up to a good size for making FFTs: --> 864 N = image.good_fft_size(N) 866 # Make sure we hit the minimum size specified in the gsparams. 867 N = max(N, self.gsparams.minimum_fft_size) File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/image.py:786](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/image.py#line=785), in Image.good_fft_size(cls, input_size) 782 # Reference from GalSim C++ 783 # https://github.com/GalSim-developers/GalSim/blob/ece3bd32c1ae6ed771f2b489c5ab1b25729e0ea4/src/Image.cpp#L1009 784 # Reduce slightly to eliminate potential rounding errors: 785 insize = (1.0 - 1.0e-5) * input_size --> 786 log2n = math.log(2.0) * math.ceil(math.log(insize) [/](https://jupyterhub.idris.fr/) math.log(2.0)) 787 log2n3 = math.log(3.0) + math.log(2.0) * math.ceil( 788 (math.log(insize) - math.log(3.0)) [/](https://jupyterhub.idris.fr/) math.log(2.0) 789 ) 790 log2n3 = max(log2n3, math.log(6.0)) # must be even number [... skipping hidden 1 frame] File [/lustre/fswork/projects/rech/ixh/ufd72rp/.local_jax_galsim/lib/python3.12/site-packages/jax/_src/core.py:1835](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/.local_jax_galsim/lib/python3.12/site-packages/jax/_src/core.py#line=1834), in concretization_function_error.<locals>.error(self, arg) 1834 def error(self, arg): -> 1835 raise ConcretizationTypeError(arg, fname_context) ConcretizationTypeError: Abstract tracer value encountered where concrete value is expected: traced array with shape float64[] The problem arose with the `float` function. If trying to convert the data type of a value, try using `x.astype(float)` or `jnp.array(x, float)` instead. The error occurred while tracing the function draw_stamp_xgalsim at [/tmp/ipykernel_760504/2509194331.py:35](https://jupyterhub.idris.fr/tmp/ipykernel_760504/2509194331.py#line=34) for jit. This concrete value was not available in Python because it depends on the values of the arguments convolved_object[0]['obj_list'][0][0]['obj_list'][0][0][0]['scale_radius'], convolved_object[0]['obj_list'][0][0]['obj_list'][0][1]['jac'], convolved_object[0]['obj_list'][0][0]['obj_list'][0][1]['offset'][0], convolved_object[0]['obj_list'][0][0]['obj_list'][0][1]['offset'][1], convolved_object[0]['obj_list'][0][0]['obj_list'][1][0][0]['nu'], convolved_object[0]['obj_list'][0][0]['obj_list'][1][0][0]['scale_radius'], convolved_object[0]['obj_list'][0][0]['obj_list'][1][1]['jac'], convolved_object[0]['obj_list'][0][0]['obj_list'][1][1]['offset'][0], convolved_object[0]['obj_list'][0][0]['obj_list'][1][1]['offset'][1], convolved_object[0]['obj_list'][1][0]['sigma'], image_pos[0], image_pos[1], and local_wcs[0]['scale']so fixing
static_argnames=("slen","fft_size")seems to not be enough.w/o the jit cf. only the vmap
draw_vmap_func = jax.vmap(partial(draw_stamp_xgalsim,convolved_object=obj,slen=52,fft_size=256)) draw_vmap_func(image_pos=image_positions, local_wcs=local_wcss)
gives
image_pos: galsim.PositionD(256.0,256.0), local_wcs: galsim.PixelScale(0.2) image_pos: galsim.PositionD(100.0,200.0), local_wcs: galsim.PixelScale(0.2) image_pos: galsim.PositionD(300.0,400.0), local_wcs: galsim.PixelScale(0.2and then a crash
-------------------------------------------------------------------------- ConcretizationTypeError Traceback (most recent call last) Cell In[32], line 1 ----> 1 draw_vmap_func(image_pos=image_positions, local_wcs=local_wcss) [... skipping hidden 6 frame] Cell In[18], line 46, in draw_stamp_xgalsim(convolved_object, image_pos, local_wcs, fft_size, slen) 43 gsparams = xgalsim.GSParams(minimum_fft_size=fft_size, maximum_fft_size=fft_size) 45 jax.debug.print("image_pos: {}, local_wcs: {}",image_pos,local_wcs) ---> 46 stamp = convolved_object.drawImage( 47 nx=slen, ny=slen, center=image_pos, wcs=local_wcs, dtype=jnp.float64 48 ) 49 return stamp File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:783](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/jax_galsim/gsobject.py#line=782), in GSObject.drawImage(self, image, nx, ny, bounds, scale, wcs, dtype, method, area, exptime, gain, add_to_image, center, use_true_center, offset, n_photons, rng, max_extra_noise, poisson_flux, sensor, photon_ops, n_subsample, maxN, save_photons, bandpass, setup_only, surface_ops) 781 added_photons = prof.drawReal(image, add_to_image) 782 else: --> 783 added_photons = prof.drawFFT(image, add_to_image) 785 image.added_flux = added_photons [/](https://jupyterhub.idris.fr/) flux_scale 786 # Restore the original center and wcs File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:928](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/jax_galsim/gsobject.py#line=927), in GSObject.drawFFT(self, image, add_to_image) 923 if image.wcs is None or not image.wcs.isPixelScale(): 924 raise _galsim.GalSimValueError( 925 "drawFFT requires an image with a PixelScale wcs", image 926 ) --> 928 kimage, wrap_size = self.drawFFT_makeKImage(image) 929 kimage = self._drawKImage(kimage) 930 return self.drawFFT_finish(image, kimage, wrap_size, add_to_image) File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:864](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/jax_galsim/gsobject.py#line=863), in GSObject.drawFFT_makeKImage(self, image) 861 N = jnp.max(jnp.array([N, image_N])) 863 # Round up to a good size for making FFTs: --> 864 N = image.good_fft_size(N) 866 # Make sure we hit the minimum size specified in the gsparams. 867 N = max(N, self.gsparams.minimum_fft_size) File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/image.py:786](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/jax_galsim/image.py#line=785), in Image.good_fft_size(cls, input_size) 782 # Reference from GalSim C++ 783 # https://github.com/GalSim-developers/GalSim/blob/ece3bd32c1ae6ed771f2b489c5ab1b25729e0ea4/src/Image.cpp#L1009 784 # Reduce slightly to eliminate potential rounding errors: 785 insize = (1.0 - 1.0e-5) * input_size --> 786 log2n = math.log(2.0) * math.ceil(math.log(insize) [/](https://jupyterhub.idris.fr/) math.log(2.0)) 787 log2n3 = math.log(3.0) + math.log(2.0) * math.ceil( 788 (math.log(insize) - math.log(3.0)) [/](https://jupyterhub.idris.fr/) math.log(2.0) 789 ) 790 log2n3 = max(log2n3, math.log(6.0)) # must be even number [... skipping hidden 1 frame] File [/lustre/fswork/projects/rech/ixh/ufd72rp/.local_jax_galsim/lib/python3.12/site-packages/jax/_src/core.py:1835](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/.local_jax_galsim/lib/python3.12/site-packages/jax/_src/core.py#line=1834), in concretization_function_error.<locals>.error(self, arg) 1834 def error(self, arg): -> 1835 raise ConcretizationTypeError(arg, fname_context) ConcretizationTypeError: Abstract tracer value encountered where concrete value is expected: traced array with shape float64[] The problem arose with the `float` function. If trying to convert the data type of a value, try using `x.astype(float)` or `jnp.array(x, float)` instead. This BatchTracer with object id 23070813019440 was created on line: [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/image.py:785:17](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/jax_galsim/image.py#line=784) (Image.good_fft_size)You have to have static params marked at every jit boundary.
I'm guessing we'll have to modify the source code.
The other thing is that idk if the variable index location of the stamp is ok for jit or not.
To recobver @ismael-mendoza "bounds" problem I have set this snippet (keeping
get_convolved_object_xgalsimabove)def draw_stamp_xgalsim( convolved_object: xgalsim.Convolution, image_pos: xgalsim.PositionD, local_wcs: xgalsim.PixelScale, *, fft_size: int, slen: int, ) -> jax.Array: gsparams = xgalsim.GSParams(minimum_fft_size=fft_size, maximum_fft_size=fft_size) convolved_object = convolved_object.withGSParams(gsparams) # <------- this force the FFT size jax.debug.print("image_pos: {}, local_wcs: {}",image_pos,local_wcs) stamp = convolved_object.drawImage( nx=slen, ny=slen, center=image_pos, wcs=local_wcs, dtype=jnp.float64 ) return stamp obj1= get_convolved_object_xgalsim(flux_d=10**3, flux_b=100, hlr_b=10, hlr_d=20, q_b=0.5, q_d=0.9, beta=jnp.pi/4) obj_list = [obj1,obj1,obj1] obj_array = jax.tree.map(lambda *vals: jnp.array(vals), *obj_list) # array-of-structs into a struct-of-arrays, image = xgalsim.Image(ncol=512, nrow=512, scale=0.2, dtype=jnp.float64) wcs = image.wcs xs=jnp.array([256,100,300]) ys=jnp.array([256,200,400]) image_positions = jax.vmap(lambda x, y: xgalsim.PositionD(x,y))(xs, ys) local_wcss = jax.vmap(lambda x: wcs.local(image_pos=x))(image_positions) # you can debug using #p_debug = jax.vmap(lambda x: jax.debug.print("{}",x)) #p_debug(obj_array) #p_debug(image_positions) #p_debug(local_wcss) draw_vmap_func = jax.vmap(partial(draw_stamp_xgalsim,slen=52,fft_size=256)) draw_vmap_func(obj_array,image_positions,local_wcss)
this will print
image_pos: galsim.PositionD(256.0,256.0), local_wcs: galsim.PixelScale(0.2)
image_pos: galsim.PositionD(100.0,200.0), local_wcs: galsim.PixelScale(0.2)
image_pos: galsim.PositionD(300.0,400.0), local_wcs: galsim.PixelScale(0.2)and then crash
--------------------------------------------------------------------------- TracerBoolConversionError Traceback (most recent call last) Cell In[106], line 1 ----> 1 draw_vmap_func(obj_array,image_positions,local_wcss) [... skipping hidden 6 frame] Cell In[91], line 12, in draw_stamp_xgalsim(convolved_object, image_pos, local_wcs, fft_size, slen) 10 convolved_object = convolved_object.withGSParams(gsparams) 11 jax.debug.print("image_pos: {}, local_wcs: {}",image_pos,local_wcs) ---> 12 stamp = convolved_object.drawImage( 13 nx=slen, ny=slen, center=image_pos, wcs=local_wcs, dtype=jnp.float64 14 ) 15 return stamp File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:783](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/gsobject.py#line=782), in GSObject.drawImage(self, image, nx, ny, bounds, scale, wcs, dtype, method, area, exptime, gain, add_to_image, center, use_true_center, offset, n_photons, rng, max_extra_noise, poisson_flux, sensor, photon_ops, n_subsample, maxN, save_photons, bandpass, setup_only, surface_ops) 781 added_photons = prof.drawReal(image, add_to_image) 782 else: --> 783 added_photons = prof.drawFFT(image, add_to_image) 785 image.added_flux = added_photons [/](https://jupyterhub.idris.fr/) flux_scale 786 # Restore the original center and wcs File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:930](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/gsobject.py#line=929), in GSObject.drawFFT(self, image, add_to_image) 928 kimage, wrap_size = self.drawFFT_makeKImage(image) 929 kimage = self._drawKImage(kimage) --> 930 return self.drawFFT_finish(image, kimage, wrap_size, add_to_image) File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/gsobject.py:913](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/gsobject.py#line=912), in GSObject.drawFFT_finish(self, image, kimage, wrap_size, add_to_image) 909 real_image = Image( 910 bounds=breal, array=real_image_arr, dtype=image.dtype, wcs=image.wcs 911 ) 912 # Add (a portion of) this to the original image. --> 913 temp = real_image.subImage(image.bounds) 914 if add_to_image: 915 image._array = image._array + temp._array File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/image.py:490](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/image.py#line=489), in Image.subImage(self, bounds) 486 if not self.bounds.isDefined(): 487 raise _galsim.GalSimUndefinedBoundsError( 488 "Attempt to access subImage of undefined image" 489 ) --> 490 if not self.bounds.includes(bounds): 491 raise _galsim.GalSimBoundsError( 492 "Attempt to access subImage not (fully) in image", bounds, self.bounds 493 ) 494 i1 = bounds.ymin - self.ymin File [/lustre/fswork/projects/rech/ixh/ufd72rp/JAX-GalSim/jax_galsim/bounds.py:112](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/JAX-GalSim/jax_galsim/bounds.py#line=111), in Bounds.includes(self, *args) 109 if isinstance(args[0], Bounds): 110 b = args[0] 111 return ( --> 112 self.isDefined() 113 and b.isDefined() 114 and self.xmin <= b.xmin 115 and self.xmax >= b.xmax 116 and self.ymin <= b.ymin 117 and self.ymax >= b.ymax 118 ) 119 elif isinstance(args[0], Position): 120 p = args[0]To focus on the "bounds pb" I have set a small snippet
from jax.tree_util import register_pytree_node_class p_debug = jax.vmap(lambda x: jax.debug.print("{}",x)) #A class that mimic de bounds usefull methods for the pb @register_pytree_node_class class A: def __init__(self, *args): self.xmin, self.xmax, self.ymin, self.ymax = args def __repr__(self): return "galsim.%s(%s,%s,%s,%s)" % ( self.__class__.__name__, self.xmin, self.xmax, self.ymin, self.ymax, ) def __str__(self): return "galsim.%s(%s,%s,%s,%s)" % ( self.__class__.__name__, self.xmin, self.xmax, self.ymin, self.ymax, ) def __hash__(self): return hash( ( self.__class__.__name__, self.xmin, self.xmax, self.ymin, self.ymax, ) ) def includes_bounds(self, *args): b = args[0] return ( self.xmin <= b.xmin and self.xmax >= b.xmax and self.ymin <= b.ymin and self.ymax >= b.ymax ) def tree_flatten(self): """This function flattens the Bounds into a list of children nodes that will be traced by JAX and auxiliary static data.""" # Define the children nodes of the PyTree that need tracing children = (self.xmin, self.xmax, self.ymin, self.ymax) # Define auxiliary static data that doesn’t need to be traced aux_data = None return (children, aux_data) @classmethod def tree_unflatten(cls, aux_data, children): """Recreates an instance of the class from flatten representation""" return cls(*children)
Then I use if like this
a0 = A(1,128,1,128) a1 = A(32,98,32,98) a2 = A(-1,10,5,200) bnd_list=[a1,a2] bnd_array = jax.tree.map(lambda *vals: jnp.array(vals), *bnd_list) # array-of-structs into a struct-of-arrays, #a debug print p_debug(bnd_array) #wiil list the 2 bounds with there parameters jax.vmap(lambda x:a0.includes_bounds(x))(bnd_array) #this will crash with the present implementation of includes_bounds
The traceback points the pb
Cell In[54], line 41, in A.includes_bounds(self, *args) 38 def includes_bounds(self, *args): 39 b = args[0] 40 return ( ---> 41 self.xmin <= b.xmin 42 and self.xmax >= b.xmax 43 and self.ymin <= b.ymin 44 and self.ymax >= b.ymax 45 ) File [/lustre/fswork/projects/rech/ixh/ufd72rp/.local_jax_galsim/lib/python3.12/site-packages/jax/_src/core.py:1829](https://jupyterhub.idris.fr/user/ufd72rp/jupyter_6/lab/tree/JAX-GalSim/.local_jax_galsim/lib/python3.12/site-packages/jax/_src/core.py#line=1828), in concretization_function_error.<locals>.error(self, arg) 1828 def error(self, arg): -> 1829 raise TracerBoolConversionError(arg) TracerBoolConversionError: Attempted boolean conversion of traced array with shape bool[]. This BatchTracer with object id 23194243169648 was created on line: [/tmp/ipykernel_979527/3062732234.py:41:12](https://jupyterhub.idris.fr/tmp/ipykernel_979527/3062732234.py#line=40) (A.includes_bounds)Ok so concerning the "bounds include" pb I think I have managed to get it right: the solution is rather simple
def includes_bounds_bis(self, b): return ( (self.xmin <= b.xmin) & (self.xmax >= b.xmax) & (self.ymin <= b.ymin) & (self.ymax >= b.ymax) )
Then,
jax.vmap(lambda x:a0.includes_bounds_bis(x))(bnd_array)
gives
Array([ True, False], dtype=bool)as expected and I can jit it
jax.jit(jax.vmap(lambda x:a0.includes_bounds_bis(x)))(bnd_array)
it gives the sames expected answer. In fact I found a statement by jakevdp
.... If you want traceable boolean logic, you need to use &, |, and ~ instead of and, or, and not ...So may be we find a way to write Bounds operation differently. Notice that it is nopt worth dto use jax.lax.cond in this scenario.
Reacted by Matthew R. Becker and Ismael MendozaI have a Google Drive here cooked a Bounds class with the modificatons and some tests vmap & jit(vmap) with Bounds, Positions, and (x,y) values
notice that the class Bounds is not complete there is no the to_galsim/from_galsim... and I have note includes the BoundsD
ismael-mendoza commented
on Mar 12, 2026 CollaboratorAuthorMore actionsThanks @beckermr and @jecampagne for looking into this.
I'm going to check my original implementation with the fixes in #209 and report back on any remaining issues.
That code won't fix everything yet FWIW.
Reacted by Ismael MendozaOK the issue here is that when you vmap over the image center, it puts tracers in the image bounds and jax does not support dynamic indexing via numpy syntax. We could try a dynamic slice operator, but in general this approach won't work unless the stamp size is fixed, which can be problematic for other reasons.
Reacted by Ismael Mendozaismael-mendoza commented
on Mar 13, 2026 CollaboratorAuthorMore actionsCould you explain why there is a problem with the stamp size being fixed? I don't think we can get around that limitation in JAX?
I guess I'm thinking that we could find a reasonable, not too large, stamp size that can accommodate most galaxies that we want to draw say in the CATSIM catalog no? (maybe there are some large galaxies that we need to handle separately)
It's not variable stamp size but the variable offset of the stamps.
Another issue is that currently the bounds are declare as auxiliary data for the pytrees and so are not traced at all.
- addedenhancementNew feature or requestNew feature or requesttestsRelated to unit-testingRelated to unit-testingand removedenhancementNew feature or requestNew feature or request
on Apr 28, 2026
for the jax galsim first paper we are interested in investigating the performance of drawing actual galaxy scenes with jax-galsim, not stamps, and compare the performance GalSim both on CPU and GPU.
We are thinking:
Sumobject directly rather than adding each object one at a time to the image, do the same in GalSim