lloca.framesnet.equi_frames.LearnedRestFrames

class lloca.framesnet.equi_frames.LearnedRestFrames(*args, compile=False, **kwargs)[source]

Bases: LearnedFrames

Rest frame transformation with learnable rotation.

This is a special case of LearnedPolarDecompositionFrames where the boost vector is chosen to be the particle momentum. Note that the rotation is constructed equivariantly to get the correct transformation behaviour of local frames.

Parameters:
  • *args – Passed to LearnedFrames

  • **kwargs – Passed to LearnedFrames

  • compile (bool) – Option to compile the orthonormalization procedure. Does not yet give significant speedups in our tests.

forward(fourmomenta, scalars=None, ptr=None, return_tracker=False, **kwargs)[source]
Parameters:
  • fourmomenta (torch.Tensor) – Tensor of shape (…, 4) containing the four-momenta

  • scalars (torch.Tensor or None) – Optional tensor of shape (…, n_scalars) containing additional scalar features

  • ptr (torch.Tensor or None) – Pointer for sparse tensors, or None for dense tensors

  • return_tracker (bool) – If True, return a tracker dictionary with regularization information

Returns:

  • Frames – Local frames constructed from the polar decomposition of the four-momenta

  • tracker (dict (optional)) – Dictionary containing regularization information, if return_tracker is True