MCPcopy Create free account
hub / github.com/NVlabs/InstantSplat / rotation2quad

Function rotation2quad

utils/pose_utils.py:117–180  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

115 return ret
116
117def 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

Callers 1

get_tensor_from_cameraFunction · 0.85

Calls 2

_sqrt_positive_partFunction · 0.85
sizeMethod · 0.80

Tested by

no test coverage detected