OpenMotionLab/MotionGPT
118
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 