Skip to content

add scripts testing full scene drawing residuals+timing #190

Description

@ismael-mendoza

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:

  • ~ 100x100 pixel images
  • use the CATSIM catalog like descwl-shear-sims, use Spergel profiles for Bulges since we don't have DeVacouleurs profile implemented.
  • use Moffat profile (need to solve moffat profile drawing is slow #161 first)
  • draw Sum object directly rather than adding each object one at a time to the image, do the same in GalSim

Activity

  1. added this to the First Paper milestone on Feb 10, 2026
  2. ismael-mendoza commented on Mar 6, 2026

    @ismael-mendoza
    CollaboratorAuthor

    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_params is 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!

  3. beckermr commented on Mar 6, 2026

    @beckermr
    Collaborator

    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 cond operator or you can add all of the objects first and then render.

  4. beckermr commented on Mar 6, 2026

    @beckermr
    Collaborator

    You should also be sure to declare the fft_size parameter as a static arg for the JIT decorator

  5. jecampagne commented on Mar 8, 2026

    @jecampagne
    Collaborator

    Hi, I have cooked this Pay attention that draw_stamp_xgalsim has 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")

    Image

    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.
    then

    local_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'].
    
  6. jecampagne commented on Mar 8, 2026

    @jecampagne
    Collaborator

    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.

  7. jecampagne commented on Mar 8, 2026

    @jecampagne
    Collaborator

    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.2
    

    and 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)
    
  8. beckermr commented on Mar 8, 2026

    @beckermr
    Collaborator

    You have to have static params marked at every jit boundary.

    I'm guessing we'll have to modify the source code.

  9. beckermr commented on Mar 8, 2026

    @beckermr
    Collaborator

    The other thing is that idk if the variable index location of the stamp is ok for jit or not.

  10. jecampagne commented on Mar 10, 2026

    @jecampagne
    Collaborator

    To recobver @ismael-mendoza "bounds" problem I have set this snippet (keeping get_convolved_object_xgalsim above)

    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]
    
  11. jecampagne commented on Mar 10, 2026

    @jecampagne
    Collaborator

    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)
    
  12. jecampagne commented on Mar 10, 2026

    @jecampagne
    Collaborator

    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.

  13. jecampagne commented on Mar 11, 2026

    @jecampagne
    Collaborator

    I 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

  14. ismael-mendoza commented on Mar 12, 2026

    @ismael-mendoza
    CollaboratorAuthor

    Thanks @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.

  15. beckermr commented on Mar 12, 2026

    @beckermr
    Collaborator

    That code won't fix everything yet FWIW.

  16. beckermr commented on Mar 12, 2026

    @beckermr
    Collaborator

    OK 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.

  17. ismael-mendoza commented on Mar 13, 2026

    @ismael-mendoza
    CollaboratorAuthor

    Could 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)

  18. beckermr commented on Mar 13, 2026

    @beckermr
    Collaborator

    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.

  19. added
    enhancementNew feature or request
    testsRelated to unit-testing
    and removed
    enhancementNew feature or request
    on Apr 28, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Labels

testsRelated to unit-testing

Type

No type

Projects

No projects

    Milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions