lloca.framesnet.equi_frames.LearnedRestFrames
- class lloca.framesnet.equi_frames.LearnedRestFrames(*args, compile=False, **kwargs)[source]
Bases:
LearnedFramesRest 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