pytc.kmatΒΆ
JAX implementation of kinetic matrix elements.
Functions
Pad |
|
Iterate (U_panel, phi_r_panel, phi_s_panel) along axis-1 of U. |
|
Calculate K1 matrix: K1_{pqrs} = sum_{i,j} w_i w_j phi_p(i) phi_q(i) nabla_i u(i, j) phi_r(j) phi_s(j) |
|
Calculate K1 kernel: K1_{kl} = sum_{g,h} w_g w_h xi_{grad}(k,g) nabla u(g,h) xi_phi(l,h) |
|
Calculate K3 matrix: K3_{pqrs} = sum_{i,j} w_i w_j phi_p(i) phi_q(i) (nabla_i u(i, j))^2 phi_r(j) phi_s(j) |
|
Calculate K3 kernel: K3_{kl} = sum_{g,h} w_g w_h xi_phi(k,g) |nabla u(g,h)|^2 xi_phi(l,h) |
|
Streaming wrapper around |
|
Contract K1 using ISDF decomposition. |
|
Streaming-capable wrapper around |
|
Compute (K1 - K2)[pqrs] in one pass, halving GPU peak vs separate calls. |
|
Streaming-capable wrapper around |
|
Contract K3 using ISDF decomposition. |
|
Streaming-capable wrapper around |