| 297 | return ci, cj, ck |
| 298 | |
| 299 | def second_derivative(self, input, ci, cj, ck): |
| 300 | u = input |
| 301 | m = u.shape[2] |
| 302 | n = u.shape[3] |
| 303 | k = u.shape[4] |
| 304 | |
| 305 | cii_0 = (u[:, :, 1, :, :] + u[:, :, 0, :, :] - |
| 306 | 2 * u[:, :, 0, :, :]).unsqueeze(2) |
| 307 | cii_1 = u[:, :, 2:, :, :] + \ |
| 308 | u[:, :, :-2, :, :] - 2 * u[:, :, 1:-1, :, :] |
| 309 | cii_2 = (u[:, :, -1, :, :] + u[:, :, -2, :, :] - |
| 310 | 2 * u[:, :, -1, :, :]).unsqueeze(2) |
| 311 | cii = torch.cat([cii_0, cii_1, cii_2], 2) |
| 312 | |
| 313 | cjj_0 = (u[:, :, :, 1, :] + u[:, :, :, 0, :] - |
| 314 | 2 * u[:, :, :, 0, :]).unsqueeze(3) |
| 315 | cjj_1 = u[:, :, :, 2:, :] + \ |
| 316 | u[:, :, :, :-2, :] - 2 * u[:, :, :, 1:-1, :] |
| 317 | cjj_2 = (u[:, :, :, -1, :] + u[:, :, :, -2, :] - |
| 318 | 2 * u[:, :, :, -1, :]).unsqueeze(3) |
| 319 | |
| 320 | cjj = torch.cat([cjj_0, cjj_1, cjj_2], 3) |
| 321 | |
| 322 | ckk_0 = (u[:, :, :, :, 1] + u[:, :, :, :, 0] - |
| 323 | 2 * u[:, :, :, :, 0]).unsqueeze(4) |
| 324 | ckk_1 = u[:, :, :, :, 2:] + \ |
| 325 | u[:, :, :, :, :-2] - 2 * u[:, :, :, :, 1:-1] |
| 326 | ckk_2 = (u[:, :, :, :, -1] + u[:, :, :, :, -2] - |
| 327 | 2 * u[:, :, :, :, -1]).unsqueeze(4) |
| 328 | |
| 329 | ckk = torch.cat([ckk_0, ckk_1, ckk_2], 4) |
| 330 | |
| 331 | cij_0 = ci[:, :, :, 1:n, :] |
| 332 | cij_1 = ci[:, :, :, -1, :].unsqueeze(3) |
| 333 | |
| 334 | cij_a = torch.cat([cij_0, cij_1], 3) |
| 335 | cij_2 = ci[:, :, :, 0, :].unsqueeze(3) |
| 336 | cij_3 = ci[:, :, :, 0:n - 1, :] |
| 337 | cij_b = torch.cat([cij_2, cij_3], 3) |
| 338 | cij = cij_a - cij_b |
| 339 | |
| 340 | cik_0 = ci[:, :, :, :, 1:n] |
| 341 | cik_1 = ci[:, :, :, :, -1].unsqueeze(4) |
| 342 | |
| 343 | cik_a = torch.cat([cik_0, cik_1], 4) |
| 344 | cik_2 = ci[:, :, :, :, 0].unsqueeze(4) |
| 345 | cik_3 = ci[:, :, :, :, 0:k - 1] |
| 346 | cik_b = torch.cat([cik_2, cik_3], 4) |
| 347 | cik = cik_a - cik_b |
| 348 | |
| 349 | cjk_0 = cj[:, :, :, :, 1:n] |
| 350 | cjk_1 = cj[:, :, :, :, -1].unsqueeze(4) |
| 351 | |
| 352 | cjk_a = torch.cat([cjk_0, cjk_1], 4) |
| 353 | cjk_2 = cj[:, :, :, :, 0].unsqueeze(4) |
| 354 | cjk_3 = cj[:, :, :, :, 0:k - 1] |
| 355 | cjk_b = torch.cat([cjk_2, cjk_3], 4) |
| 356 | cjk = cjk_a - cjk_b |