|
1 | | -import jax.numpy as jnp |
2 | 1 | import numpy as np |
3 | 2 |
|
4 | 3 | from typing import List, Optional, Tuple |
5 | 4 |
|
6 | | -from autoconf import cached_property |
7 | | - |
8 | 5 | from autoarray import type as ty |
9 | 6 | from autoarray.inversion.linear_obj.neighbors import Neighbors |
10 | 7 | from autoarray.mask.mask_2d import Mask2D |
@@ -89,18 +86,18 @@ def overlay_grid( |
89 | 86 | """ |
90 | 87 | grid = grid.array |
91 | 88 |
|
92 | | - y_min = jnp.min(grid[:, 0]) - buffer |
93 | | - y_max = jnp.max(grid[:, 0]) + buffer |
94 | | - x_min = jnp.min(grid[:, 1]) - buffer |
95 | | - x_max = jnp.max(grid[:, 1]) + buffer |
| 89 | + y_min = xp.min(grid[:, 0]) - buffer |
| 90 | + y_max = xp.max(grid[:, 0]) + buffer |
| 91 | + x_min = xp.min(grid[:, 1]) - buffer |
| 92 | + x_max = xp.max(grid[:, 1]) + buffer |
96 | 93 |
|
97 | | - pixel_scales = jnp.array( |
| 94 | + pixel_scales = xp.array( |
98 | 95 | ( |
99 | 96 | (y_max - y_min) / shape_native[0], |
100 | 97 | (x_max - x_min) / shape_native[1], |
101 | 98 | ) |
102 | 99 | ) |
103 | | - origin = jnp.array(((y_max + y_min) / 2.0, (x_max + x_min) / 2.0)) |
| 100 | + origin = xp.array(((y_max + y_min) / 2.0, (x_max + x_min) / 2.0)) |
104 | 101 |
|
105 | 102 | grid_slim = grid_2d_util.grid_2d_slim_via_shape_native_not_mask_from( |
106 | 103 | shape_native=shape_native, |
|
0 commit comments