Team Ai
Apppublic

OpenMotionLab/MotionGPT

sourceHugging Facemitupdated 1y agoView on Hugging Face
118likes
geometry_tools.py567 linesDownload Raw Back to utils
1# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved.2# Check PYTORCH3D_LICENCE before use3 4import functools5from typing import Optional6 7import torch8import torch.nn.functional as F9 10 11"""12The transformation matrices returned from the functions in this file assume13the points on which the transformation will be applied are column vectors.14i.e. the R matrix is structured as15 16    R = [17            [Rxx, Rxy, Rxz],18            [Ryx, Ryy, Ryz],19            [Rzx, Rzy, Rzz],20        ]  # (3, 3)21 22This matrix can be applied to column vectors by post multiplication23by the points e.g.24 25    points = [[0], [1], [2]]  # (3 x 1) xyz coordinates of a point26    transformed_points = R * points27 28To apply the same matrix to points which are row vectors, the R matrix29can be transposed and pre multiplied by the points:30 31e.g.32    points = [[0, 1, 2]]  # (1 x 3) xyz coordinates of a point33    transformed_points = points * R.transpose(1, 0)34"""35 36 37# Added38def matrix_of_angles(cos, sin, inv=False, dim=2):39    assert dim in [2, 3]40    sin = -sin if inv else sin41    if dim == 2:42        row1 = torch.stack((cos, -sin), axis=-1)43        row2 = torch.stack((sin, cos), axis=-1)44        return torch.stack((row1, row2), axis=-2)45    elif dim == 3:46        row1 = torch.stack((cos, -sin, 0*cos), axis=-1)47        row2 = torch.stack((sin, cos, 0*cos), axis=-1)48        row3 = torch.stack((0*sin, 0*cos, 1+0*cos), axis=-1)49        return torch.stack((row1, row2, row3),axis=-2)50 51 52def quaternion_to_matrix(quaternions):53    """54    Convert rotations given as quaternions to rotation matrices.55 56    Args:57        quaternions: quaternions with real part first,58            as tensor of shape (..., 4).59 60    Returns:61        Rotation matrices as tensor of shape (..., 3, 3).62    """63    r, i, j, k = torch.unbind(quaternions, -1)64    two_s = 2.0 / (quaternions * quaternions).sum(-1)65 66    o = torch.stack(67        (68            1 - two_s * (j * j + k * k),69            two_s * (i * j - k * r),70            two_s * (i * k + j * r),71            two_s * (i * j + k * r),72            1 - two_s * (i * i + k * k),73            two_s * (j * k - i * r),74            two_s * (i * k - j * r),75            two_s * (j * k + i * r),76            1 - two_s * (i * i + j * j),77        ),78        -1,79    )80    return o.reshape(quaternions.shape[:-1] + (3, 3))81 82 83def _copysign(a, b):84    """85    Return a tensor where each element has the absolute value taken from the,86    corresponding element of a, with sign taken from the corresponding87    element of b. This is like the standard copysign floating-point operation,88    but is not careful about negative 0 and NaN.89 90    Args:91        a: source tensor.92        b: tensor whose signs will be used, of the same shape as a.93 94    Returns:95        Tensor of the same shape as a with the signs of b.96    """97    signs_differ = (a < 0) != (b < 0)98    return torch.where(signs_differ, -a, a)99 100 101def _sqrt_positive_part(x):102    """103    Returns torch.sqrt(torch.max(0, x))104    but with a zero subgradient where x is 0.105    """106    ret = torch.zeros_like(x)107    positive_mask = x > 0108    ret[positive_mask] = torch.sqrt(x[positive_mask])109    return ret110 111 112def matrix_to_quaternion(matrix):113    """114    Convert rotations given as rotation matrices to quaternions.115 116    Args:117        matrix: Rotation matrices as tensor of shape (..., 3, 3).118 119    Returns:120        quaternions with real part first, as tensor of shape (..., 4).121    """122    if matrix.size(-1) != 3 or matrix.size(-2) != 3:123        raise ValueError(f"Invalid rotation matrix  shape f{matrix.shape}.")124    m00 = matrix[..., 0, 0]125    m11 = matrix[..., 1, 1]126    m22 = matrix[..., 2, 2]127    o0 = 0.5 * _sqrt_positive_part(1 + m00 + m11 + m22)128    x = 0.5 * _sqrt_positive_part(1 + m00 - m11 - m22)129    y = 0.5 * _sqrt_positive_part(1 - m00 + m11 - m22)130    z = 0.5 * _sqrt_positive_part(1 - m00 - m11 + m22)131    o1 = _copysign(x, matrix[..., 2, 1] - matrix[..., 1, 2])132    o2 = _copysign(y, matrix[..., 0, 2] - matrix[..., 2, 0])133    o3 = _copysign(z, matrix[..., 1, 0] - matrix[..., 0, 1])134    return torch.stack((o0, o1, o2, o3), -1)135 136 137def _axis_angle_rotation(axis: str, angle):138    """139    Return the rotation matrices for one of the rotations about an axis140    of which Euler angles describe, for each value of the angle given.141 142    Args:143        axis: Axis label "X" or "Y or "Z".144        angle: any shape tensor of Euler angles in radians145 146    Returns:147        Rotation matrices as tensor of shape (..., 3, 3).148    """149 150    cos = torch.cos(angle)151    sin = torch.sin(angle)152    one = torch.ones_like(angle)153    zero = torch.zeros_like(angle)154 155    if axis == "X":156        R_flat = (one, zero, zero, zero, cos, -sin, zero, sin, cos)157    if axis == "Y":158        R_flat = (cos, zero, sin, zero, one, zero, -sin, zero, cos)159    if axis == "Z":160        R_flat = (cos, -sin, zero, sin, cos, zero, zero, zero, one)161 162    return torch.stack(R_flat, -1).reshape(angle.shape + (3, 3))163 164 165def euler_angles_to_matrix(euler_angles, convention: str):166    """167    Convert rotations given as Euler angles in radians to rotation matrices.168 169    Args:170        euler_angles: Euler angles in radians as tensor of shape (..., 3).171        convention: Convention string of three uppercase letters from172            {"X", "Y", and "Z"}.173 174    Returns:175        Rotation matrices as tensor of shape (..., 3, 3).176    """177    if euler_angles.dim() == 0 or euler_angles.shape[-1] != 3:178        raise ValueError("Invalid input euler angles.")179    if len(convention) != 3:180        raise ValueError("Convention must have 3 letters.")181    if convention[1] in (convention[0], convention[2]):182        raise ValueError(f"Invalid convention {convention}.")183    for letter in convention:184        if letter not in ("X", "Y", "Z"):185            raise ValueError(f"Invalid letter {letter} in convention string.")186    matrices = map(_axis_angle_rotation, convention, torch.unbind(euler_angles, -1))187    return functools.reduce(torch.matmul, matrices)188 189 190def _angle_from_tan(191    axis: str, other_axis: str, data, horizontal: bool, tait_bryan: bool192):193    """194    Extract the first or third Euler angle from the two members of195    the matrix which are positive constant times its sine and cosine.196 197    Args:198        axis: Axis label "X" or "Y or "Z" for the angle we are finding.199        other_axis: Axis label "X" or "Y or "Z" for the middle axis in the200            convention.201        data: Rotation matrices as tensor of shape (..., 3, 3).202        horizontal: Whether we are looking for the angle for the third axis,203            which means the relevant entries are in the same row of the204            rotation matrix. If not, they are in the same column.205        tait_bryan: Whether the first and third axes in the convention differ.206 207    Returns:208        Euler Angles in radians for each matrix in data as a tensor209        of shape (...).210    """211 212    i1, i2 = {"X": (2, 1), "Y": (0, 2), "Z": (1, 0)}[axis]213    if horizontal:214        i2, i1 = i1, i2215    even = (axis + other_axis) in ["XY", "YZ", "ZX"]216    if horizontal == even:217        return torch.atan2(data[..., i1], data[..., i2])218    if tait_bryan:219        return torch.atan2(-data[..., i2], data[..., i1])220    return torch.atan2(data[..., i2], -data[..., i1])221 222 223def _index_from_letter(letter: str):224    if letter == "X":225        return 0226    if letter == "Y":227        return 1228    if letter == "Z":229        return 2230 231 232def matrix_to_euler_angles(matrix, convention: str):233    """234    Convert rotations given as rotation matrices to Euler angles in radians.235 236    Args:237        matrix: Rotation matrices as tensor of shape (..., 3, 3).238        convention: Convention string of three uppercase letters.239 240    Returns:241        Euler angles in radians as tensor of shape (..., 3).242    """243    if len(convention) != 3:244        raise ValueError("Convention must have 3 letters.")245    if convention[1] in (convention[0], convention[2]):246        raise ValueError(f"Invalid convention {convention}.")247    for letter in convention:248        if letter not in ("X", "Y", "Z"):249            raise ValueError(f"Invalid letter {letter} in convention string.")250    if matrix.size(-1) != 3 or matrix.size(-2) != 3:251        raise ValueError(f"Invalid rotation matrix  shape f{matrix.shape}.")252    i0 = _index_from_letter(convention[0])253    i2 = _index_from_letter(convention[2])254    tait_bryan = i0 != i2255    if tait_bryan:256        central_angle = torch.asin(257            matrix[..., i0, i2] * (-1.0 if i0 - i2 in [-1, 2] else 1.0)258        )259    else:260        central_angle = torch.acos(matrix[..., i0, i0])261 262    o = (263        _angle_from_tan(264            convention[0], convention[1], matrix[..., i2], False, tait_bryan265        ),266        central_angle,267        _angle_from_tan(268            convention[2], convention[1], matrix[..., i0, :], True, tait_bryan269        ),270    )271    return torch.stack(o, -1)272 273 274def random_quaternions(275    n: int, dtype: Optional[torch.dtype] = None, device=None, requires_grad=False276):277    """278    Generate random quaternions representing rotations,279    i.e. versors with nonnegative real part.280 281    Args:282        n: Number of quaternions in a batch to return.283        dtype: Type to return.284        device: Desired device of returned tensor. Default:285            uses the current device for the default tensor type.286        requires_grad: Whether the resulting tensor should have the gradient287            flag set.288 289    Returns:290        Quaternions as tensor of shape (N, 4).291    """292    o = torch.randn((n, 4), dtype=dtype, device=device, requires_grad=requires_grad)293    s = (o * o).sum(1)294    o = o / _copysign(torch.sqrt(s), o[:, 0])[:, None]295    return o296 297 298def random_rotations(299    n: int, dtype: Optional[torch.dtype] = None, device=None, requires_grad=False300):301    """302    Generate random rotations as 3x3 rotation matrices.303 304    Args:305        n: Number of rotation matrices in a batch to return.306        dtype: Type to return.307        device: Device of returned tensor. Default: if None,308            uses the current device for the default tensor type.309        requires_grad: Whether the resulting tensor should have the gradient310            flag set.311 312    Returns:313        Rotation matrices as tensor of shape (n, 3, 3).314    """315    quaternions = random_quaternions(316        n, dtype=dtype, device=device, requires_grad=requires_grad317    )318    return quaternion_to_matrix(quaternions)319 320 321def random_rotation(322    dtype: Optional[torch.dtype] = None, device=None, requires_grad=False323):324    """325    Generate a single random 3x3 rotation matrix.326 327    Args:328        dtype: Type to return329        device: Device of returned tensor. Default: if None,330            uses the current device for the default tensor type331        requires_grad: Whether the resulting tensor should have the gradient332            flag set333 334    Returns:335        Rotation matrix as tensor of shape (3, 3).336    """337    return random_rotations(1, dtype, device, requires_grad)[0]338 339 340def standardize_quaternion(quaternions):341    """342    Convert a unit quaternion to a standard form: one in which the real343    part is non negative.344 345    Args:346        quaternions: Quaternions with real part first,347            as tensor of shape (..., 4).348 349    Returns:350        Standardized quaternions as tensor of shape (..., 4).351    """352    return torch.where(quaternions[..., 0:1] < 0, -quaternions, quaternions)353 354 355def quaternion_raw_multiply(a, b):356    """357    Multiply two quaternions.358    Usual torch rules for broadcasting apply.359 360    Args:361        a: Quaternions as tensor of shape (..., 4), real part first.362        b: Quaternions as tensor of shape (..., 4), real part first.363 364    Returns:365        The product of a and b, a tensor of quaternions shape (..., 4).366    """367    aw, ax, ay, az = torch.unbind(a, -1)368    bw, bx, by, bz = torch.unbind(b, -1)369    ow = aw * bw - ax * bx - ay * by - az * bz370    ox = aw * bx + ax * bw + ay * bz - az * by371    oy = aw * by - ax * bz + ay * bw + az * bx372    oz = aw * bz + ax * by - ay * bx + az * bw373    return torch.stack((ow, ox, oy, oz), -1)374 375 376def quaternion_multiply(a, b):377    """378    Multiply two quaternions representing rotations, returning the quaternion379    representing their composition, i.e. the versor with nonnegative real part.380    Usual torch rules for broadcasting apply.381 382    Args:383        a: Quaternions as tensor of shape (..., 4), real part first.384        b: Quaternions as tensor of shape (..., 4), real part first.385 386    Returns:387        The product of a and b, a tensor of quaternions of shape (..., 4).388    """389    ab = quaternion_raw_multiply(a, b)390    return standardize_quaternion(ab)391 392 393def quaternion_invert(quaternion):394    """395    Given a quaternion representing rotation, get the quaternion representing396    its inverse.397 398    Args:399        quaternion: Quaternions as tensor of shape (..., 4), with real part400            first, which must be versors (unit quaternions).401 402    Returns:403        The inverse, a tensor of quaternions of shape (..., 4).404    """405 406    return quaternion * quaternion.new_tensor([1, -1, -1, -1])407 408 409def quaternion_apply(quaternion, point):410    """411    Apply the rotation given by a quaternion to a 3D point.412    Usual torch rules for broadcasting apply.413 414    Args:415        quaternion: Tensor of quaternions, real part first, of shape (..., 4).416        point: Tensor of 3D points of shape (..., 3).417 418    Returns:419        Tensor of rotated points of shape (..., 3).420    """421    if point.size(-1) != 3:422        raise ValueError(f"Points are not in 3D, f{point.shape}.")423    real_parts = point.new_zeros(point.shape[:-1] + (1,))424    point_as_quaternion = torch.cat((real_parts, point), -1)425    out = quaternion_raw_multiply(426        quaternion_raw_multiply(quaternion, point_as_quaternion),427        quaternion_invert(quaternion),428    )429    return out[..., 1:]430 431 432def axis_angle_to_matrix(axis_angle):433    """434    Convert rotations given as axis/angle to rotation matrices.435 436    Args:437        axis_angle: Rotations given as a vector in axis angle form,438            as a tensor of shape (..., 3), where the magnitude is439            the angle turned anticlockwise in radians around the440            vector's direction.441 442    Returns:443        Rotation matrices as tensor of shape (..., 3, 3).444    """445    return quaternion_to_matrix(axis_angle_to_quaternion(axis_angle))446 447 448def matrix_to_axis_angle(matrix):449    """450    Convert rotations given as rotation matrices to axis/angle.451 452    Args:453        matrix: Rotation matrices as tensor of shape (..., 3, 3).454 455    Returns:456        Rotations given as a vector in axis angle form, as a tensor457            of shape (..., 3), where the magnitude is the angle458            turned anticlockwise in radians around the vector's459            direction.460    """461    return quaternion_to_axis_angle(matrix_to_quaternion(matrix))462 463 464def axis_angle_to_quaternion(axis_angle):465    """466    Convert rotations given as axis/angle to quaternions.467 468    Args:469        axis_angle: Rotations given as a vector in axis angle form,470            as a tensor of shape (..., 3), where the magnitude is471            the angle turned anticlockwise in radians around the472            vector's direction.473 474    Returns:475        quaternions with real part first, as tensor of shape (..., 4).476    """477    angles = torch.norm(axis_angle, p=2, dim=-1, keepdim=True)478    half_angles = 0.5 * angles479    eps = 1e-6480    small_angles = angles.abs() < eps481    sin_half_angles_over_angles = torch.empty_like(angles)482    sin_half_angles_over_angles[~small_angles] = (483        torch.sin(half_angles[~small_angles]) / angles[~small_angles]484    )485    # for x small, sin(x/2) is about x/2 - (x/2)^3/6486    # so sin(x/2)/x is about 1/2 - (x*x)/48487    sin_half_angles_over_angles[small_angles] = (488        0.5 - (angles[small_angles] * angles[small_angles]) / 48489    )490    quaternions = torch.cat(491        [torch.cos(half_angles), axis_angle * sin_half_angles_over_angles], dim=-1492    )493    return quaternions494 495 496def quaternion_to_axis_angle(quaternions):497    """498    Convert rotations given as quaternions to axis/angle.499 500    Args:501        quaternions: quaternions with real part first,502            as tensor of shape (..., 4).503 504    Returns:505        Rotations given as a vector in axis angle form, as a tensor506            of shape (..., 3), where the magnitude is the angle507            turned anticlockwise in radians around the vector's508            direction.509    """510    norms = torch.norm(quaternions[..., 1:], p=2, dim=-1, keepdim=True)511    half_angles = torch.atan2(norms, quaternions[..., :1])512    angles = 2 * half_angles513    eps = 1e-6514    small_angles = angles.abs() < eps515    sin_half_angles_over_angles = torch.empty_like(angles)516    sin_half_angles_over_angles[~small_angles] = (517        torch.sin(half_angles[~small_angles]) / angles[~small_angles]518    )519    # for x small, sin(x/2) is about x/2 - (x/2)^3/6520    # so sin(x/2)/x is about 1/2 - (x*x)/48521    sin_half_angles_over_angles[small_angles] = (522        0.5 - (angles[small_angles] * angles[small_angles]) / 48523    )524    return quaternions[..., 1:] / sin_half_angles_over_angles525 526 527def rotation_6d_to_matrix(d6: torch.Tensor) -> torch.Tensor:528    """529    Converts 6D rotation representation by Zhou et al. [1] to rotation matrix530    using Gram--Schmidt orthogonalisation per Section B of [1].531    Args:532        d6: 6D rotation representation, of size (*, 6)533 534    Returns:535        batch of rotation matrices of size (*, 3, 3)536 537    [1] Zhou, Y., Barnes, C., Lu, J., Yang, J., & Li, H.538    On the Continuity of Rotation Representations in Neural Networks.539    IEEE Conference on Computer Vision and Pattern Recognition, 2019.540    Retrieved from http://arxiv.org/abs/1812.07035541    """542 543    a1, a2 = d6[..., :3], d6[..., 3:]544    b1 = F.normalize(a1, dim=-1)545    b2 = a2 - (b1 * a2).sum(-1, keepdim=True) * b1546    b2 = F.normalize(b2, dim=-1)547    b3 = torch.cross(b1, b2, dim=-1)548    return torch.stack((b1, b2, b3), dim=-2)549 550 551def matrix_to_rotation_6d(matrix: torch.Tensor) -> torch.Tensor:552    """553    Converts rotation matrices to 6D rotation representation by Zhou et al. [1]554    by dropping the last row. Note that 6D representation is not unique.555    Args:556        matrix: batch of rotation matrices of size (*, 3, 3)557 558    Returns:559        6D rotation representation, of size (*, 6)560 561    [1] Zhou, Y., Barnes, C., Lu, J., Yang, J., & Li, H.562    On the Continuity of Rotation Representations in Neural Networks.563    IEEE Conference on Computer Vision and Pattern Recognition, 2019.564    Retrieved from http://arxiv.org/abs/1812.07035565    """566    return matrix[..., :2, :].clone().reshape(*matrix.size()[:-2], 6)567