Skip to content

Commit 44adda3

Browse files
authored
Merge pull request #150 from rcjackson/smoothness_gradients
FIX: Correct axis parameters and remove duplicate coefficient scaling…
2 parents 2c75e81 + f7c7d3f commit 44adda3

2 files changed

Lines changed: 19 additions & 97 deletions

File tree

pydda/cost_functions/_cost_functions_jax.py

Lines changed: 12 additions & 90 deletions
Original file line numberDiff line numberDiff line change
@@ -191,26 +191,26 @@ def calculate_smoothness_cost(u, v, w, dx, dy, dz, Cx=1e-5, Cy=1e-5, Cz=1e-5):
191191
Cx
192192
* (
193193
jnp.gradient(dudx, dx, axis=2)
194-
+ jnp.gradient(dvdx, dx, axis=1)
194+
+ jnp.gradient(dvdx, dx, axis=2)
195195
+ jnp.gradient(dwdx, dx, axis=2)
196196
)
197197
** 2
198198
)
199199
y_term = (
200200
Cy
201201
* (
202-
jnp.gradient(dudy, dy, axis=2)
202+
jnp.gradient(dudy, dy, axis=1)
203203
+ jnp.gradient(dvdy, dy, axis=1)
204-
+ jnp.gradient(dwdy, dy, axis=2)
204+
+ jnp.gradient(dwdy, dy, axis=1)
205205
)
206206
** 2
207207
)
208208
z_term = (
209209
Cz
210210
* (
211-
jnp.gradient(dudz, dz, axis=2)
212-
+ jnp.gradient(dvdz, dz, axis=1)
213-
+ jnp.gradient(dwdz, dz, axis=2)
211+
jnp.gradient(dudz, dz, axis=0)
212+
+ jnp.gradient(dvdz, dz, axis=0)
213+
+ jnp.gradient(dwdz, dz, axis=0)
214214
)
215215
** 2
216216
)
@@ -253,94 +253,16 @@ def calculate_smoothness_gradient(
253253
y: float array
254254
value of gradient of smoothness cost function
255255
"""
256-
dudx = jnp.gradient(u, dx, axis=2)
257-
dudy = jnp.gradient(u, dy, axis=1)
258-
dudz = jnp.gradient(u, dz, axis=0)
259-
dvdx = jnp.gradient(v, dx, axis=2)
260-
dvdy = jnp.gradient(v, dy, axis=1)
261-
dvdz = jnp.gradient(v, dz, axis=0)
262-
dwdx = jnp.gradient(w, dx, axis=2)
263-
dwdy = jnp.gradient(w, dy, axis=1)
264-
dwdz = jnp.gradient(w, dz, axis=0)
265-
266-
x_term = (
267-
Cx
268-
* (
269-
jnp.gradient(dudx, dx, axis=2)
270-
+ jnp.gradient(dvdx, dx, axis=1)
271-
+ jnp.gradient(dwdx, dx, axis=2)
272-
)
273-
** 2
274-
)
275-
y_term = (
276-
Cy
277-
* (
278-
jnp.gradient(dudy, dy, axis=2)
279-
+ jnp.gradient(dvdy, dy, axis=1)
280-
+ jnp.gradient(dwdy, dy, axis=2)
281-
)
282-
** 2
283-
)
284-
z_term = (
285-
Cz
286-
* (
287-
jnp.gradient(dudz, dz, axis=2)
288-
+ jnp.gradient(dvdz, dz, axis=1)
289-
+ jnp.gradient(dwdz, dz, axis=2)
290-
)
291-
** 2
292-
)
293-
294-
du = x_term / dx
295-
dv = y_term / dy
296-
dw = z_term / dz
297-
dudx = jnp.gradient(du, dx, axis=2)
298-
dudy = jnp.gradient(du, dy, axis=1)
299-
dudz = jnp.gradient(du, dz, axis=0)
300-
dvdx = jnp.gradient(dv, dx, axis=2)
301-
dvdy = jnp.gradient(dv, dy, axis=1)
302-
dvdz = jnp.gradient(dv, dz, axis=0)
303-
dwdx = jnp.gradient(dw, dx, axis=2)
304-
dwdy = jnp.gradient(dw, dy, axis=1)
305-
dwdz = jnp.gradient(dw, dz, axis=0)
306-
307-
x_term = (
308-
Cx
309-
* (
310-
jnp.gradient(dudx, dx, axis=2)
311-
+ jnp.gradient(dvdx, dx, axis=1)
312-
+ jnp.gradient(dwdx, dx, axis=2)
313-
)
314-
** 2
315-
)
316-
y_term = (
317-
Cy
318-
* (
319-
jnp.gradient(dudy, dy, axis=2)
320-
+ jnp.gradient(dvdy, dy, axis=1)
321-
+ jnp.gradient(dwdy, dy, axis=2)
322-
)
323-
** 2
324-
)
325-
z_term = (
326-
Cz
327-
* (
328-
jnp.gradient(dudz, dz, axis=2)
329-
+ jnp.gradient(dvdz, dz, axis=1)
330-
+ jnp.gradient(dwdz, dz, axis=2)
331-
)
332-
** 2
256+
primals, fun_vjp = jax.vjp(
257+
calculate_smoothness_cost, u, v, w, dx, dy, dz, Cx, Cy, Cz
333258
)
334-
335-
grad_u = x_term / dx
336-
grad_v = y_term / dy
337-
grad_w = z_term / dz
259+
grad_u, grad_v, grad_w, _, _, _, _, _, _ = fun_vjp(1.0)
338260

339261
# Impermeability condition
340-
grad_w.at[0, :, :].set(0)
262+
grad_w = grad_w.at[0, :, :].set(0)
341263
if upper_bc is True:
342-
grad_w.at[-1, :, :].set(0)
343-
y = jnp.stack([grad_u * Cx * 2, grad_v * Cy * 2, grad_w * Cz * 2], axis=0)
264+
grad_w = grad_w.at[-1, :, :].set(0)
265+
y = jnp.stack([grad_u, grad_v, grad_w], axis=0)
344266

345267
return y.flatten()
346268

pydda/cost_functions/_cost_functions_numpy.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -183,26 +183,26 @@ def calculate_smoothness_cost(u, v, w, dx, dy, dz, Cx=1e-5, Cy=1e-5, Cz=1e-5):
183183
Cx
184184
* (
185185
np.gradient(dudx, dx, axis=2)
186-
+ np.gradient(dvdx, dx, axis=1)
186+
+ np.gradient(dvdx, dx, axis=2)
187187
+ np.gradient(dwdx, dx, axis=2)
188188
)
189189
** 2
190190
)
191191
y_term = (
192192
Cy
193193
* (
194-
np.gradient(dudy, dy, axis=2)
194+
np.gradient(dudy, dy, axis=1)
195195
+ np.gradient(dvdy, dy, axis=1)
196-
+ np.gradient(dwdy, dy, axis=2)
196+
+ np.gradient(dwdy, dy, axis=1)
197197
)
198198
** 2
199199
)
200200
z_term = (
201201
Cz
202202
* (
203-
np.gradient(dudz, dz, axis=2)
204-
+ np.gradient(dvdz, dz, axis=1)
205-
+ np.gradient(dwdz, dz, axis=2)
203+
np.gradient(dudz, dz, axis=0)
204+
+ np.gradient(dvdz, dz, axis=0)
205+
+ np.gradient(dwdz, dz, axis=0)
206206
)
207207
** 2
208208
)
@@ -260,7 +260,7 @@ def calculate_smoothness_gradient(
260260
if upper_bc is True:
261261
grad_w[-1, :, :] = 0
262262

263-
y = np.stack([grad_u * Cx * 2, grad_v * Cy * 2, grad_w * Cz * 2], axis=0)
263+
y = np.stack([grad_u, grad_v, grad_w], axis=0)
264264

265265
return y.flatten()
266266

0 commit comments

Comments
 (0)