assemble_tt_variation_jets_trs#

t3toolbox.backend.sampling_derivatives.assemble_tt_variation_jets_trs(sigma_tildes, tau_tildes, deta_tildes, xi_jets, mu_jets, nu_jets, trs, n_probe, sum_over_probes)#
def assemble_tt_variation_jets_trs(
        sigma_tildes:   typ.Sequence[NDArray],  # len=d, elm_shape=(order+1,)+W+K+C+(rR(i+1),)
        tau_tildes:     typ.Sequence[NDArray],  # len=d, elm_shape=(order+1,)+W+K+C+(rL(i+1),)
        deta_tildes:    typ.Sequence[NDArray],  # len=d, elm_shape=(order+1,)+W+K+C+(nUi,)
        xi_jets:        typ.Sequence[NDArray],  # frame input jets, len=d, elm_shape=(2,)+W+C+(nUi,)
        mu_jets:        typ.Sequence[NDArray],  # len=d, elm_shape=(order+1,)+W+C+(rLi,)
        nu_jets:        typ.Sequence[NDArray],  # len=d, elm_shape=(order+1,)+W+C+(rR(i+1),)
        trs:            NDArray,                # binomial tensor, shape=(order+1,order+1,order+1)
        n_probe:        int,                    # number of sample-stack (W) axes
        sum_over_probes: bool,
) -> typ.Tuple[NDArray, ...]:                   # dG_tildes. len=d, elm_shape=[W+]C+(rLi,nUi,rRi)

Assemble TT-core variation gradients (the 3-edge, trs case): three order-less trs outer products mu (x) xi (x) sigma_tilde + tau_tilde (x) xi (x) nu + mu (x) deta_tilde (x) nu (the core-adjoints of the forward sigma / tau / deta contractions).

Parameters:
  • sigma_tildes (t3toolbox.backend.common.typ.Sequence[NDArray])

  • tau_tildes (t3toolbox.backend.common.typ.Sequence[NDArray])

  • deta_tildes (t3toolbox.backend.common.typ.Sequence[NDArray])

  • xi_jets (t3toolbox.backend.common.typ.Sequence[NDArray])

  • mu_jets (t3toolbox.backend.common.typ.Sequence[NDArray])

  • nu_jets (t3toolbox.backend.common.typ.Sequence[NDArray])

  • trs (NDArray)

  • n_probe (int)

  • sum_over_probes (bool)

Return type:

t3toolbox.backend.common.typ.Tuple[NDArray, Ellipsis]