Convert rotations given as rotation matrices to quaternions. Args: matrix: Rotation matrices as tensor of shape (..., 3, 3). Returns: quaternions with real part first, as tensor of shape (..., 4). Source: https://pytorch3d.readthedocs.io/en/latest/_modules/pytorch3
(matrix: torch.Tensor)
| 115 | return ret |
| 116 | |
| 117 | def rotation2quad(matrix: torch.Tensor) -> torch.Tensor: |
| 118 | """ |
| 119 | Convert rotations given as rotation matrices to quaternions. |
| 120 | |
| 121 | Args: |
| 122 | matrix: Rotation matrices as tensor of shape (..., 3, 3). |
| 123 | |
| 124 | Returns: |
| 125 | quaternions with real part first, as tensor of shape (..., 4). |
| 126 | Source: https://pytorch3d.readthedocs.io/en/latest/_modules/pytorch3d/transforms/rotation_conversions.html#matrix_to_quaternion |
| 127 | """ |
| 128 | if matrix.size(-1) != 3 or matrix.size(-2) != 3: |
| 129 | raise ValueError(f"Invalid rotation matrix shape {matrix.shape}.") |
| 130 | |
| 131 | if not isinstance(matrix, torch.Tensor): |
| 132 | matrix = torch.tensor(matrix).cuda() |
| 133 | |
| 134 | batch_dim = matrix.shape[:-2] |
| 135 | m00, m01, m02, m10, m11, m12, m20, m21, m22 = torch.unbind( |
| 136 | matrix.reshape(batch_dim + (9,)), dim=-1 |
| 137 | ) |
| 138 | |
| 139 | q_abs = _sqrt_positive_part( |
| 140 | torch.stack( |
| 141 | [ |
| 142 | 1.0 + m00 + m11 + m22, |
| 143 | 1.0 + m00 - m11 - m22, |
| 144 | 1.0 - m00 + m11 - m22, |
| 145 | 1.0 - m00 - m11 + m22, |
| 146 | ], |
| 147 | dim=-1, |
| 148 | ) |
| 149 | ) |
| 150 | |
| 151 | # we produce the desired quaternion multiplied by each of r, i, j, k |
| 152 | quat_by_rijk = torch.stack( |
| 153 | [ |
| 154 | # pyre-fixme[58]: `**` is not supported for operand types `Tensor` and |
| 155 | # `int`. |
| 156 | torch.stack([q_abs[..., 0] ** 2, m21 - m12, m02 - m20, m10 - m01], dim=-1), |
| 157 | # pyre-fixme[58]: `**` is not supported for operand types `Tensor` and |
| 158 | # `int`. |
| 159 | torch.stack([m21 - m12, q_abs[..., 1] ** 2, m10 + m01, m02 + m20], dim=-1), |
| 160 | # pyre-fixme[58]: `**` is not supported for operand types `Tensor` and |
| 161 | # `int`. |
| 162 | torch.stack([m02 - m20, m10 + m01, q_abs[..., 2] ** 2, m12 + m21], dim=-1), |
| 163 | # pyre-fixme[58]: `**` is not supported for operand types `Tensor` and |
| 164 | # `int`. |
| 165 | torch.stack([m10 - m01, m20 + m02, m21 + m12, q_abs[..., 3] ** 2], dim=-1), |
| 166 | ], |
| 167 | dim=-2, |
| 168 | ) |
| 169 | |
| 170 | # We floor here at 0.1 but the exact level is not important; if q_abs is small, |
| 171 | # the candidate won't be picked. |
| 172 | flr = torch.tensor(0.1).to(dtype=q_abs.dtype, device=q_abs.device) |
| 173 | quat_candidates = quat_by_rijk / (2.0 * q_abs[..., None].max(flr)) |
| 174 |
no test coverage detected