lloca.reps.tensorreps_transform.TensorRepsTransform
- class lloca.reps.tensorreps_transform.TensorRepsTransform(reps)[source]
Bases:
ModuleTensor representation transformation module.
- Parameters:
reps (TensorReps) – Tensor representations to transform, sorted by order.
- forward(tensor, frames)[source]
Apply a transformation to a tensor of a given representation.
- Parameters:
tensor (torch.Tensor) – The tensor to transform, shape (…, self.reps.dim).
frames (Frames) – The local frames to apply the transformation with, shape (…, 4, 4).
- Returns:
The transformed tensor, shape (…, self.reps.dim).
- Return type:
torch.Tensor