@@ -144,9 +144,11 @@ def traced_grid_2d_list_from(
144144 redshift_list = [galaxies [0 ].redshift for galaxies in planes ]
145145
146146 for plane_index , galaxies in enumerate (planes ):
147- scaled_grid = grid .copy ()
147+
148+ scaled_grid = grid .array
148149
149150 if plane_index > 0 :
151+
150152 for previous_plane_index in range (plane_index ):
151153 scaling_factor = cosmology .scaling_factor_between_redshifts_from (
152154 redshift_0 = redshift_list [previous_plane_index ],
@@ -156,10 +158,14 @@ def traced_grid_2d_list_from(
156158 )
157159
158160 scaled_deflections = (
159- scaling_factor * traced_deflection_list [previous_plane_index ]
161+ scaling_factor * traced_deflection_list [previous_plane_index ]. array
160162 )
161163
162- scaled_grid -= scaled_deflections
164+ scaled_grid = scaled_grid - scaled_deflections
165+
166+ scaled_grid = aa .Grid2DIrregular (
167+ values = scaled_grid ,
168+ )
163169
164170 traced_grid_list .append (scaled_grid )
165171
@@ -168,12 +174,7 @@ def traced_grid_2d_list_from(
168174 return traced_grid_list
169175
170176 deflections_yx_2d = sum (
171- map (lambda g : g .deflections_yx_2d_from (grid = scaled_grid , xp = xp ), galaxies )
172- )
173-
174- # Remove NaN deflection values to sanitize the ray-tracing calculation for JAX.
175- deflections_yx_2d = xp .where (
176- xp .isfinite (deflections_yx_2d .array ), deflections_yx_2d .array , 0.0
177+ (g .deflections_yx_2d_from (grid = scaled_grid , xp = xp ) for g in galaxies )
177178 )
178179
179180 traced_deflection_list .append (deflections_yx_2d )
0 commit comments