lloca.reps.tensorreps_transform.TensorRepsTransform

class lloca.reps.tensorreps_transform.TensorRepsTransform(reps)[source]

Bases: Module

Tensor 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