@@ -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
0 commit comments