FastQuat 1.0b2__tar.gz → 1.0b4__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (30) hide show
  1. fastquat-1.0b4/.github/workflows/ci.yml +71 -0
  2. {fastquat-1.0b2 → fastquat-1.0b4}/.pre-commit-config.yaml +6 -14
  3. {fastquat-1.0b2 → fastquat-1.0b4}/PKG-INFO +4 -2
  4. {fastquat-1.0b2 → fastquat-1.0b4}/README.md +3 -1
  5. {fastquat-1.0b2 → fastquat-1.0b4}/docs/source/api/quaternion.md +1 -0
  6. {fastquat-1.0b2 → fastquat-1.0b4}/docs/source/development.md +17 -20
  7. {fastquat-1.0b2 → fastquat-1.0b4}/docs/source/index.md +1 -0
  8. {fastquat-1.0b2 → fastquat-1.0b4}/docs/source/user-guide/installation.md +2 -2
  9. {fastquat-1.0b2 → fastquat-1.0b4}/pyproject.toml +11 -0
  10. {fastquat-1.0b2 → fastquat-1.0b4}/src/fastquat/quaternion.py +82 -57
  11. {fastquat-1.0b2 → fastquat-1.0b4}/tests/test_pow.py +48 -0
  12. {fastquat-1.0b2 → fastquat-1.0b4}/tests/test_rotation.py +74 -0
  13. {fastquat-1.0b2 → fastquat-1.0b4}/uv.lock +22 -0
  14. fastquat-1.0b2/.github/workflows/ci.yml +0 -67
  15. {fastquat-1.0b2 → fastquat-1.0b4}/.github/workflows/release.yml +0 -0
  16. {fastquat-1.0b2 → fastquat-1.0b4}/.gitignore +0 -0
  17. {fastquat-1.0b2 → fastquat-1.0b4}/.readthedocs.yaml +0 -0
  18. {fastquat-1.0b2 → fastquat-1.0b4}/docs/.gitignore +0 -0
  19. {fastquat-1.0b2 → fastquat-1.0b4}/docs/Makefile +0 -0
  20. {fastquat-1.0b2 → fastquat-1.0b4}/docs/README.md +0 -0
  21. {fastquat-1.0b2 → fastquat-1.0b4}/docs/source/conf.py +0 -0
  22. {fastquat-1.0b2 → fastquat-1.0b4}/docs/source/user-guide/getting-started.ipynb +0 -0
  23. {fastquat-1.0b2 → fastquat-1.0b4}/docs/source/user-guide/tutorial-rotations.ipynb +0 -0
  24. {fastquat-1.0b2 → fastquat-1.0b4}/docs/source/user-guide/tutorial-slerp.ipynb +0 -0
  25. {fastquat-1.0b2 → fastquat-1.0b4}/src/fastquat/__init__.py +0 -0
  26. {fastquat-1.0b2 → fastquat-1.0b4}/tests/conftest.py +0 -0
  27. {fastquat-1.0b2 → fastquat-1.0b4}/tests/test_base.py +0 -0
  28. {fastquat-1.0b2 → fastquat-1.0b4}/tests/test_indexing.py +0 -0
  29. {fastquat-1.0b2 → fastquat-1.0b4}/tests/test_math.py +0 -0
  30. {fastquat-1.0b2 → fastquat-1.0b4}/tests/test_tensor.py +0 -0
@@ -0,0 +1,71 @@
1
+ name: CI
2
+
3
+ permissions: {}
4
+
5
+ on:
6
+ push:
7
+ branches:
8
+ - main
9
+ pull_request:
10
+ workflow_dispatch:
11
+
12
+ concurrency:
13
+ group: ${{ github.workflow }}-${{ github.ref_name }}-${{ github.event.pull_request.number || github.sha }}
14
+ cancel-in-progress: true
15
+
16
+ jobs:
17
+
18
+ pre-commit:
19
+ name: Pre-commit
20
+ runs-on: ubuntu-latest
21
+ timeout-minutes: 5
22
+ steps:
23
+ - uses: actions/checkout@v7.0.1
24
+ - name: Install uv
25
+ uses: astral-sh/setup-uv@v10.0.1
26
+ with:
27
+ version: '0.12.18'
28
+ - name: Cache prek
29
+ uses: actions/cache@v6.1.0
30
+ with:
31
+ path: ~/.cache/prek
32
+ key: prek-${{ hashFiles('.pre-commit-config.yaml') }}
33
+ restore-keys: |
34
+ prek-
35
+ # ty resolves third-party imports from the environment, so jax has to be installed.
36
+ - name: Install dependencies
37
+ run: uv sync --locked --no-default-groups --group lint
38
+ - name: Run pre-commit
39
+ run: uv run --no-sync prek run --all-files --show-diff-on-failure --color always
40
+
41
+ test:
42
+ name: Run tests on Python ${{ matrix.python-version }}
43
+ runs-on: ubuntu-latest
44
+ timeout-minutes: 15
45
+ strategy:
46
+ fail-fast: false
47
+ matrix:
48
+ python-version: ['3.10', '3.11', '3.12', '3.13', '3.14']
49
+
50
+ steps:
51
+ - name: Checkout code
52
+ uses: actions/checkout@v7.0.1
53
+ with:
54
+ fetch-depth: 0
55
+
56
+ - name: Install uv
57
+ uses: astral-sh/setup-uv@v10.0.1
58
+ with:
59
+ version: '0.12.18'
60
+ python-version: ${{ matrix.python-version }}
61
+
62
+ - name: Build wheel
63
+ run: uv build --wheel
64
+
65
+ - name: Install project's wheel and dependencies
66
+ run: |
67
+ uv sync --locked --no-install-project
68
+ uv pip install dist/fastquat-*.whl
69
+
70
+ - name: Run tests
71
+ run: uv run --no-sync pytest -v
@@ -7,9 +7,9 @@ repos:
7
7
  - --all
8
8
 
9
9
  - repo: https://github.com/astral-sh/ruff-pre-commit
10
- rev: "v0.14.7"
10
+ rev: "v0.16.9"
11
11
  hooks:
12
- - id: ruff
12
+ - id: ruff-check
13
13
  name: ruff linting
14
14
  - id: ruff-format
15
15
  name: ruff formatting
@@ -31,21 +31,13 @@ repos:
31
31
  - id: check-merge-conflict
32
32
 
33
33
  - repo: https://github.com/kynan/nbstripout
34
- rev: 0.6.1
34
+ rev: "0.9.1"
35
35
  hooks:
36
36
  - id: nbstripout
37
37
  name: notebook stripping
38
38
  args: [--extra-keys=metadata.language_info.version]
39
39
 
40
- - repo: https://github.com/pre-commit/mirrors-mypy
41
- rev: 'v1.19.0'
40
+ - repo: https://github.com/astral-sh/ty-pre-commit
41
+ rev: v0.0.84
42
42
  hooks:
43
- - id: mypy
44
- additional_dependencies:
45
- - jax
46
- args:
47
- - --strict
48
- - --show-error-codes
49
- - --enable-error-code=ignore-without-code
50
- - --allow-untyped-calls
51
- files: ^fastquat/
43
+ - id: ty
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: FastQuat
3
- Version: 1.0b2
3
+ Version: 1.0b4
4
4
  Summary: High-performance quaternions with JAX support
5
5
  Project-URL: homepage, https://fastquat.readthedocs.io
6
6
  Project-URL: repository, https://github.com/CMBSciPol/fastquat
@@ -89,7 +89,7 @@ interpolated = q1.slerp(q2, t=0.5) # Halfway between q1 and q2
89
89
  ### Core Operations
90
90
  - **Quaternion arithmetic**: Addition, multiplication, conjugation, inverse, power, exponentiation, logarithm
91
91
  - **Normalization**: Efficient unit quaternion computation
92
- - **Conversion**: To/from rotation matrices, Euler angles
92
+ - **Conversion**: To/from rotation matrices, axis-angle, Euler angles
93
93
  - **Vector rotation**: Direct vector transformation
94
94
 
95
95
  ### Advanced Interpolation
@@ -124,11 +124,13 @@ from fastquat import Quaternion
124
124
  key = jax.random.PRNGKey(42)
125
125
  q_batch = Quaternion.random(key, shape=(1000,))
126
126
 
127
+
127
128
  # JIT-compiled batch operations
128
129
  @jax.jit
129
130
  def batch_rotate(quaternions, vectors):
130
131
  return quaternions.rotate_vector(vectors)
131
132
 
133
+
132
134
  vectors = jax.random.normal(key, (1000, 3))
133
135
  rotated_batch = batch_rotate(q_batch, vectors)
134
136
  ```
@@ -66,7 +66,7 @@ interpolated = q1.slerp(q2, t=0.5) # Halfway between q1 and q2
66
66
  ### Core Operations
67
67
  - **Quaternion arithmetic**: Addition, multiplication, conjugation, inverse, power, exponentiation, logarithm
68
68
  - **Normalization**: Efficient unit quaternion computation
69
- - **Conversion**: To/from rotation matrices, Euler angles
69
+ - **Conversion**: To/from rotation matrices, axis-angle, Euler angles
70
70
  - **Vector rotation**: Direct vector transformation
71
71
 
72
72
  ### Advanced Interpolation
@@ -101,11 +101,13 @@ from fastquat import Quaternion
101
101
  key = jax.random.PRNGKey(42)
102
102
  q_batch = Quaternion.random(key, shape=(1000,))
103
103
 
104
+
104
105
  # JIT-compiled batch operations
105
106
  @jax.jit
106
107
  def batch_rotate(quaternions, vectors):
107
108
  return quaternions.rotate_vector(vectors)
108
109
 
110
+
109
111
  vectors = jax.random.normal(key, (1000, 3))
110
112
  rotated_batch = batch_rotate(q_batch, vectors)
111
113
  ```
@@ -18,6 +18,7 @@ compilation, automatic differentiation, and vectorization.
18
18
  .. automethod:: fastquat.Quaternion.from_array
19
19
  .. automethod:: fastquat.Quaternion.from_scalar_vector
20
20
  .. automethod:: fastquat.Quaternion.from_rotation_matrix
21
+ .. automethod:: fastquat.Quaternion.from_axis_angle
21
22
  .. automethod:: fastquat.Quaternion.zeros
22
23
  .. automethod:: fastquat.Quaternion.ones
23
24
  .. automethod:: fastquat.Quaternion.full
@@ -112,20 +112,21 @@ from jax.typing import ArrayLike
112
112
 
113
113
 
114
114
  class Quaternion:
115
- ...
116
- def new_method(self, parameter: ArrayLike) -> Self:
117
- """Brief description of what the method does.
118
-
119
- Args:
120
- parameter: Description of the parameter
121
-
122
- Returns:
123
- Description of the return value
124
- """
125
- # Implementation using JAX operations
126
- parameter = jnp.asarray(parameter)
127
- result = jnp.some_operation(self.wxyz, parameter)
128
- return Quaternion.from_array(result)
115
+ ...
116
+
117
+ def new_method(self, parameter: ArrayLike) -> Self:
118
+ """Brief description of what the method does.
119
+
120
+ Args:
121
+ parameter: Description of the parameter
122
+
123
+ Returns:
124
+ Description of the return value
125
+ """
126
+ # Implementation using JAX operations
127
+ parameter = jnp.asarray(parameter)
128
+ result = jnp.some_operation(self.wxyz, parameter)
129
+ return Quaternion.from_array(result)
129
130
  ```
130
131
 
131
132
  Then add tests:
@@ -138,15 +139,11 @@ import pytest
138
139
  from fastquat import Quaternion
139
140
 
140
141
 
141
- @pytest.mark.parametrize(
142
- 'parameter, expected_values',
143
- [
144
- (..., ...),
145
- ]
146
- )
142
+ @pytest.mark.parametrize('parameter, expected_values', [(..., ...)])
147
143
  @pytest.mark.parametrize('do_jit', [False, True])
148
144
  def test_new_method(parameter, expected_values, do_jit):
149
145
  """Test the new method."""
146
+
150
147
  def test_fn(q_, parameter_):
151
148
  return q_.new_method(parameter_)
152
149
 
@@ -39,6 +39,7 @@ q2 = Quaternion(0.7071, 0.7071, 0.0, 0.0) # 90° rotation around x-axis
39
39
  | Normalization | `q.normalize()` | Unit quaternion |
40
40
  | Conjugate | `q.conj()` | Quaternion conjugate |
41
41
  | Rotation | `q.rotate_vector(v)` | Rotate 3D vector |
42
+ | Axis-angle | `Quaternion.from_axis_angle(axis, angle)` | Rotation about an axis |
42
43
  | SLERP | `q1.slerp(q2, t)` | Spherical interpolation |
43
44
  | Log | `q.log()` | Quaternion logarithm |
44
45
  | Exp | `q.exp()` | Quaternion exponential |
@@ -52,12 +52,12 @@ from fastquat import Quaternion
52
52
 
53
53
  # Create a simple quaternion
54
54
  q = Quaternion(1.0)
55
- print(f"Identity quaternion: {q}")
55
+ print(f'Identity quaternion: {q}')
56
56
 
57
57
  # Test SLERP functionality
58
58
  q2 = Quaternion(0.7071, 0.7071, 0.0, 0.0)
59
59
  interpolated = q.slerp(q2, 0.5)
60
- print(f"SLERP result: {interpolated}")
60
+ print(f'SLERP result: {interpolated}')
61
61
  ```
62
62
 
63
63
  If this runs without errors, FastQuat is properly installed!
@@ -34,6 +34,7 @@ repository = 'https://github.com/CMBSciPol/fastquat'
34
34
  [dependency-groups]
35
35
  dev = [
36
36
  {include-group = 'docs'},
37
+ {include-group = 'lint'},
37
38
  'pytest>=9.0',
38
39
  'pytest-notebook>=0.10',
39
40
  'visu-hlo',
@@ -48,6 +49,7 @@ docs = [
48
49
  'numpy>=1.24.0',
49
50
  'pillow>=9.0.0',
50
51
  ]
52
+ lint = ['prek>=0.5.4']
51
53
  cuda12 = ['jax[cuda12]>=0.4']
52
54
  cuda13 = ['jax[cuda13]>=0.7; python_version>="3.11"']
53
55
 
@@ -83,5 +85,14 @@ select = [
83
85
  [tool.ruff.format]
84
86
  quote-style = 'single'
85
87
 
88
+ [tool.ty.rules]
89
+ dynamic-function-decorator-return = "error"
90
+ missing-type-argument = "error"
91
+ possibly-unresolved-reference = "warn"
92
+ unused-ignore-comment = "error"
93
+
94
+ [tool.ty.src]
95
+ include = ["src"]
96
+
86
97
  [tool.uv]
87
98
  cache-keys = [{ git = true }]
@@ -118,6 +118,32 @@ class Quaternion:
118
118
 
119
119
  return cls.from_array(jnp.stack([w, x, y, z], axis=-1))
120
120
 
121
+ @classmethod
122
+ def from_axis_angle(cls, axis: ArrayLike, angle: ArrayLike) -> Self:
123
+ """Create the unit quaternion of a rotation about an axis.
124
+
125
+ The rotation follows the right-hand rule: a positive angle rotates counterclockwise when
126
+ looking from the tip of the axis towards the origin.
127
+
128
+ Args:
129
+ axis: Array of shape (..., 3) for the rotation axis. It does not need to be normalized.
130
+ angle: Array of shape (...) for the rotation angle, in radians.
131
+
132
+ Returns:
133
+ Quaternion of shape broadcast(axis.shape[:-1], angle.shape).
134
+ """
135
+ axis = jnp.asarray(axis)
136
+ angle = jnp.asarray(angle)
137
+ if axis.shape[-1:] != (3,):
138
+ raise ValueError(f'Axis must have shape (..., 3), got {axis.shape}')
139
+ dtype = jnp.result_type(axis, angle, float)
140
+ unit_axis = axis / jnp.linalg.norm(axis, axis=-1, keepdims=True)
141
+ half_angle = 0.5 * angle.astype(dtype)
142
+ scalar = jnp.cos(half_angle)
143
+ vector = jnp.sin(half_angle)[..., None] * unit_axis.astype(dtype)
144
+ scalar, vector = jnp.broadcast_arrays(scalar[..., None], vector)
145
+ return cls.from_scalar_vector(scalar[..., 0], vector)
146
+
121
147
  @classmethod
122
148
  def zeros(cls, shape: tuple[int, ...], dtype: DTypeLike | None = None) -> Self:
123
149
  """Create quaternions with all components set to 0.
@@ -167,7 +193,7 @@ class Quaternion:
167
193
 
168
194
  @classmethod
169
195
  def random(
170
- cls, key: jax.random.PRNGKey, shape: tuple[int, ...] = (), dtype: DTypeLike | None = None
196
+ cls, key: Array, shape: tuple[int, ...] = (), dtype: DTypeLike | None = None
171
197
  ) -> Self:
172
198
  """Generate normalized random quaternions.
173
199
 
@@ -180,7 +206,7 @@ class Quaternion:
180
206
  Normalized Quaternion.
181
207
  """
182
208
  data = jax.random.normal(key, shape + (4,), dtype=dtype)
183
- return Quaternion.from_array(data).normalize()
209
+ return cls.from_array(data).normalize()
184
210
 
185
211
  @property
186
212
  def w(self) -> Array:
@@ -214,13 +240,13 @@ class Quaternion:
214
240
  returns the quaternion [NaN, NaN, NaN, NaN].
215
241
  """
216
242
  norm = abs(self)
217
- return Quaternion.from_array(self.wxyz / jnp.expand_dims(norm, axis=-1))
243
+ return self.from_array(self.wxyz / jnp.expand_dims(norm, axis=-1))
218
244
 
219
245
  def _inverse(self) -> Self:
220
246
  """Quaternion inverse (private method - use 1/q instead)."""
221
247
  conj = self.conj()
222
248
  norm_sq = jnp.sum(self.wxyz**2, axis=-1)
223
- return Quaternion.from_array(conj.wxyz / jnp.expand_dims(norm_sq, axis=-1))
249
+ return self.from_array(conj.wxyz / jnp.expand_dims(norm_sq, axis=-1))
224
250
 
225
251
  def to_components(self) -> tuple[Array, Array, Array, Array]:
226
252
  return self.w, self.x, self.y, self.z
@@ -290,13 +316,13 @@ class Quaternion:
290
316
  if self.ndim == 0:
291
317
  raise TypeError('iteration over a 0-d quaternion')
292
318
  for i in range(self.shape[0]):
293
- yield Quaternion.from_array(self.wxyz[i])
319
+ yield self.from_array(self.wxyz[i])
294
320
 
295
321
  def __getitem__(self, idx: Any) -> Self:
296
322
  """Index or slice the tensor of quaternions."""
297
323
  if not isinstance(idx, tuple):
298
324
  idx = (idx,)
299
- return Quaternion.from_array(self.wxyz[(*idx, slice(None))])
325
+ return self.from_array(self.wxyz[(*idx, slice(None))])
300
326
 
301
327
  def __pos__(self) -> Self:
302
328
  """Quaternion positive."""
@@ -304,12 +330,12 @@ class Quaternion:
304
330
 
305
331
  def __neg__(self) -> Self:
306
332
  """Quaternion negation."""
307
- return Quaternion.from_array(-self.wxyz)
333
+ return self.from_array(-self.wxyz)
308
334
 
309
335
  def __add__(self, other: Any) -> Self:
310
336
  """Quaternion addition."""
311
337
  if isinstance(other, Quaternion):
312
- return Quaternion.from_array(self.wxyz + other.wxyz)
338
+ return self.from_array(self.wxyz + other.wxyz)
313
339
 
314
340
  try:
315
341
  other = jnp.asarray(other)
@@ -319,7 +345,7 @@ class Quaternion:
319
345
  if jnp.iscomplexobj(other):
320
346
  raise NotImplementedError('Quaternion and complex addition is not implemented.')
321
347
 
322
- return Quaternion.from_scalar_vector(self.w + other, self.vector)
348
+ return self.from_scalar_vector(self.w + other, self.vector)
323
349
 
324
350
  def __radd__(self, other: Any) -> Self:
325
351
  """Quaternion addition."""
@@ -328,7 +354,7 @@ class Quaternion:
328
354
  def __sub__(self, other: Any) -> Self:
329
355
  """Quaternion subtraction."""
330
356
  if isinstance(other, Quaternion):
331
- return Quaternion.from_array(self.wxyz - other.wxyz)
357
+ return self.from_array(self.wxyz - other.wxyz)
332
358
 
333
359
  try:
334
360
  other = jnp.asarray(other)
@@ -338,7 +364,7 @@ class Quaternion:
338
364
  if jnp.iscomplexobj(other):
339
365
  raise NotImplementedError('Quaternion and complex subtraction is not implemented.')
340
366
 
341
- return Quaternion.from_scalar_vector(self.w - other, self.vector)
367
+ return self.from_scalar_vector(self.w - other, self.vector)
342
368
 
343
369
  def __rsub__(self, other: Any) -> Self:
344
370
  """Quaternion subtraction."""
@@ -350,7 +376,7 @@ class Quaternion:
350
376
  if jnp.iscomplexobj(other):
351
377
  raise NotImplementedError('Quaternion and complex subtraction is not implemented.')
352
378
 
353
- return Quaternion.from_scalar_vector(other - self.w, -self.vector)
379
+ return self.from_scalar_vector(other - self.w, -self.vector)
354
380
 
355
381
  def __mul__(self, other: Any) -> Self:
356
382
  """Quaternion multiplication."""
@@ -363,7 +389,7 @@ class Quaternion:
363
389
  y = w1 * y2 - x1 * z2 + y1 * w2 + z1 * x2
364
390
  z = w1 * z2 + x1 * y2 - y1 * x2 + z1 * w2
365
391
 
366
- return Quaternion(w, x, y, z)
392
+ return self.from_array(jnp.stack([w, x, y, z], axis=-1))
367
393
 
368
394
  try:
369
395
  other = jnp.asarray(other)
@@ -373,7 +399,7 @@ class Quaternion:
373
399
  if jnp.iscomplexobj(other):
374
400
  raise NotImplementedError('Quaternion and complex multiplication is not implemented.')
375
401
 
376
- return Quaternion.from_array(self.wxyz * jnp.expand_dims(other, axis=-1))
402
+ return self.from_array(self.wxyz * jnp.expand_dims(other, axis=-1))
377
403
 
378
404
  def __rmul__(self, other: Any) -> Self:
379
405
  """Quaternion multiplication."""
@@ -385,7 +411,7 @@ class Quaternion:
385
411
  if jnp.iscomplexobj(other):
386
412
  raise NotImplementedError('Quaternion and complex multiplication is not implemented.')
387
413
 
388
- return Quaternion.from_array(jnp.expand_dims(other, axis=-1) * self.wxyz)
414
+ return self.from_array(jnp.expand_dims(other, axis=-1) * self.wxyz)
389
415
 
390
416
  def __truediv__(self, other: Any) -> Self:
391
417
  """Quaternion division."""
@@ -400,7 +426,7 @@ class Quaternion:
400
426
  if jnp.iscomplexobj(other):
401
427
  raise NotImplementedError('Quaternion and complex division is not implemented.')
402
428
 
403
- return Quaternion.from_array(self.wxyz / jnp.expand_dims(other, axis=-1))
429
+ return self.from_array(self.wxyz / jnp.expand_dims(other, axis=-1))
404
430
 
405
431
  def __rtruediv__(self, other: Any) -> Self:
406
432
  """Quaternion division."""
@@ -437,7 +463,7 @@ class Quaternion:
437
463
  elif exponent == -1:
438
464
  return self._inverse()
439
465
  elif exponent == 0:
440
- return Quaternion.ones(self.shape, self.dtype)
466
+ return self.ones(self.shape, self.dtype)
441
467
  elif exponent == 1:
442
468
  return self
443
469
  elif exponent == 2:
@@ -445,8 +471,9 @@ class Quaternion:
445
471
  return (exponent * self.log()).exp()
446
472
 
447
473
  # General case: q^n = exp(n * log(q))
474
+ exponent = jnp.asarray(exponent)
448
475
  result = (exponent * self.log()).exp().wxyz
449
- return Quaternion.from_array(
476
+ return self.from_array(
450
477
  jnp.where(
451
478
  exponent[..., None] == 0, jnp.array([1.0, 0.0, 0.0, 0.0], dtype=self.dtype), result
452
479
  )
@@ -455,63 +482,61 @@ class Quaternion:
455
482
  def log(self) -> Self:
456
483
  """Compute quaternion logarithm.
457
484
 
458
- For a quaternion q = ‖q‖ * (cos(θ) + sin(θ)v), the logarithm is:
459
- log(q) = log(‖q‖) + θ * v
485
+ For a quaternion q = |q| * (cos(θ) + sin(θ)v), the logarithm is:
486
+ log(q) = log(|q|) + θ * v
460
487
 
461
488
  For the zero quaternion, returns (-inf, 0, 0, 0).
462
489
 
463
490
  Returns:
464
491
  The logarithm of the quaternion
465
492
  """
466
- q_norm = abs(self)
467
-
468
- # Normalize manually to handle zero quaternion (returns 0 instead of NaN)
469
- safe_norm = jnp.where(q_norm == 0, 1.0, q_norm)
470
- unit_wxyz = self.wxyz / jnp.expand_dims(safe_norm, axis=-1)
471
-
472
- # For unit quaternion q = cos(θ) + sin(θ)v, compute θ and v
473
- # θ = arccos(w) and v = vector/|vector|
474
- unit_w = unit_wxyz[..., 0]
475
- unit_vector = unit_wxyz[..., 1:]
476
- theta = jnp.arccos(jnp.clip(unit_w, -1.0, 1.0))
477
- vector_norm = jnp.linalg.norm(unit_vector, axis=-1)
478
-
479
- # Handle case where vector is zero (real quaternion)
480
- inv_vector_norm = jnp.where(vector_norm == 0, 0.0, 1 / vector_norm)
481
- unit_vector = unit_vector * inv_vector_norm[..., None]
482
-
483
- # log(q) = log(|q|) + θ * v
484
- log_norm = jnp.log(q_norm)
485
- log_q_vector = theta[..., None] * unit_vector
493
+ scalar_part = self.w
494
+ vector_part = self.vector
495
+ vector_norm_sq = jnp.sum(vector_part**2, axis=-1)
496
+ is_real = vector_norm_sq == 0
497
+
498
+ # log(q) = log(|q|) + θ * v/|v|, with θ = atan2(|v|, s).
499
+ # The where guards keep the gradients finite for real quaternions (|v| = 0).
500
+ log_norm = 0.5 * jnp.log(scalar_part**2 + vector_norm_sq)
501
+ safe_vector_norm = jnp.sqrt(jnp.where(is_real, 1.0, vector_norm_sq))
502
+ safe_scalar_part = jnp.where(scalar_part == 0, 1.0, scalar_part)
503
+ # θ/|v| tends to 1/s when |v| → 0 (s > 0)
504
+ theta_over_vector_norm = jnp.where(
505
+ is_real,
506
+ 1 / safe_scalar_part,
507
+ jnp.arctan2(safe_vector_norm, scalar_part) / safe_vector_norm,
508
+ )
509
+ log_q_vector = theta_over_vector_norm[..., None] * vector_part
486
510
 
487
- return Quaternion.from_scalar_vector(log_norm, log_q_vector)
511
+ return self.from_scalar_vector(log_norm, log_q_vector)
488
512
 
489
513
  def exp(self) -> Self:
490
514
  """Compute quaternion exponential.
491
515
 
492
516
  For a quaternion q = s + v, the exponential is:
493
- exp(q) = exp(s) * (cos(‖v‖) + sin(‖v‖) * v/‖v‖)
517
+ exp(q) = exp(s) * (cos(|v|) + sin(|v|) * v/|v|)
494
518
 
495
519
  Returns:
496
520
  The exponential of the quaternion
497
521
  """
498
522
  scalar_part = self.w
499
523
  vector_part = self.vector
500
- vector_norm = jnp.linalg.norm(vector_part, axis=-1)
524
+ vector_norm_sq = jnp.sum(vector_part**2, axis=-1)
525
+ is_real = vector_norm_sq == 0
501
526
 
502
- # exp(s + v) = exp(s) * (cos(|v|) + sin(|v|) * v/|v|)
527
+ # The where guard keeps the gradients finite for real quaternions (|v| = 0).
528
+ vector_norm = jnp.where(is_real, 0.0, jnp.sqrt(jnp.where(is_real, 1.0, vector_norm_sq)))
503
529
  exp_scalar = jnp.exp(scalar_part)
504
- cos_vnorm = jnp.cos(vector_norm)
505
- sin_vnorm = jnp.sin(vector_norm)
506
-
507
- # Handle case where |v| = 0 (real quaternion)
508
- inv_vector_norm = jnp.where(vector_norm == 0, 0.0, 1 / vector_norm)
509
- unit_v = vector_part * jnp.expand_dims(inv_vector_norm, -1)
530
+ # sin(|v|)/|v| and cos(|v|) = 1 - |v|²/2 (sin(|v|/2)/(|v|/2))², written with sinc
531
+ # so that the first and second derivatives are exact at |v| = 0
532
+ sinc_vnorm = jnp.sinc(vector_norm / jnp.pi)
533
+ sinc_half_vnorm = jnp.sinc(vector_norm / (2 * jnp.pi))
534
+ cos_vnorm = 1 - 0.5 * vector_norm_sq * sinc_half_vnorm**2
510
535
 
511
536
  result_w = exp_scalar * cos_vnorm
512
- result_vector = exp_scalar * jnp.expand_dims(sin_vnorm, -1) * unit_v
537
+ result_vector = jnp.expand_dims(exp_scalar * sinc_vnorm, -1) * vector_part
513
538
 
514
- return Quaternion.from_scalar_vector(result_w, result_vector)
539
+ return self.from_scalar_vector(result_w, result_vector)
515
540
 
516
541
  @property
517
542
  def nbytes(self) -> int:
@@ -564,12 +589,12 @@ class Quaternion:
564
589
 
565
590
  def squeeze(self, axis=None) -> Self:
566
591
  """Supprime les dimensions de taille 1"""
567
- return Quaternion.from_array(jnp.squeeze(self.wxyz, axis=axis))
592
+ return self.from_array(jnp.squeeze(self.wxyz, axis=axis))
568
593
 
569
594
  def conjugate(self) -> Self:
570
595
  """Quaternion conjugate."""
571
596
  sign = jnp.array([1, -1, -1, -1])
572
- return Quaternion.from_array(self.wxyz * sign)
597
+ return self.from_array(self.wxyz * sign)
573
598
 
574
599
  def conj(self) -> Self:
575
600
  """Quaternion conjugate."""
@@ -617,7 +642,7 @@ class Quaternion:
617
642
 
618
643
  # Linear interpolation case
619
644
  result_linear = q1.wxyz + jnp.expand_dims(t * (1 - t), -1) * (q2_corrected - q1.wxyz)
620
- result_linear = Quaternion.from_array(result_linear).normalize()
645
+ result_linear = self.from_array(result_linear).normalize()
621
646
 
622
647
  # Spherical interpolation case
623
648
  theta = jnp.arccos(jnp.clip(dot, 0.0, 1.0))
@@ -632,9 +657,9 @@ class Quaternion:
632
657
  result_slerp = (
633
658
  jnp.expand_dims(factor1, -1) * q1.wxyz + jnp.expand_dims(factor2, -1) * q2_corrected
634
659
  )
635
- result_slerp = Quaternion.from_array(result_slerp)
660
+ result_slerp = self.from_array(result_slerp)
636
661
 
637
662
  # Choose between linear and spherical interpolation
638
663
  result = jnp.where(jnp.expand_dims(use_linear, -1), result_linear.wxyz, result_slerp.wxyz)
639
664
 
640
- return Quaternion.from_array(result)
665
+ return self.from_array(result)
@@ -1,5 +1,6 @@
1
1
  import jax
2
2
  import jax.numpy as jnp
3
+ import jax.test_util
3
4
  import pytest
4
5
 
5
6
  from fastquat import Quaternion
@@ -386,3 +387,50 @@ def test_pow_zero_quaternion(cast_type, do_jit):
386
387
  # 0^n for n > 0 should give 0
387
388
  result_0_2 = func(zero_q, cast_type(2.0))
388
389
  assert jnp.allclose(result_0_2.wxyz, zero_q.wxyz)
390
+
391
+
392
+ # Gradients of logarithm and exponential
393
+ @pytest.mark.parametrize(
394
+ 'wxyz',
395
+ [
396
+ [0.0, 0.0, 0.0, 0.0], # pure zero: exp at the identity
397
+ [0.3, 0.0, 0.0, 0.0], # real
398
+ [0.0, 0.1, -0.2, 0.3], # pure imaginary
399
+ [0.5, 0.1, -0.2, 0.3],
400
+ ],
401
+ )
402
+ def test_exp_grads(wxyz, enable_x64):
403
+ """Test exp derivatives up to second order, including at a zero vector part."""
404
+ jax.test_util.check_grads(
405
+ lambda a: Quaternion.from_array(a).exp().wxyz, (jnp.array(wxyz),), order=2
406
+ )
407
+
408
+
409
+ @pytest.mark.parametrize(
410
+ 'wxyz',
411
+ [
412
+ [1.0, 0.0, 0.0, 0.0], # identity
413
+ [2.0, 0.0, 0.0, 0.0], # real
414
+ [0.0, 0.1, -0.2, 0.3], # pure imaginary
415
+ [0.5, 0.1, -0.2, 0.3],
416
+ ],
417
+ )
418
+ def test_log_grads(wxyz, enable_x64):
419
+ """Test log derivatives, including at a zero vector part."""
420
+ jax.test_util.check_grads(
421
+ lambda a: Quaternion.from_array(a).log().wxyz, (jnp.array(wxyz),), order=1
422
+ )
423
+
424
+
425
+ def test_exp_jacobian_at_zero():
426
+ """Test that d exp(v)/dv at v = 0 maps the vector part onto itself."""
427
+ jac = jax.jacobian(lambda v: Quaternion.from_scalar_vector(0.0, v).exp().wxyz)(jnp.zeros(3))
428
+ expected = jnp.concatenate([jnp.zeros((1, 3)), jnp.eye(3)])
429
+ assert jnp.allclose(jac, expected)
430
+
431
+
432
+ def test_log_jacobian_at_identity():
433
+ """Test that d log(1 + v)/dv at v = 0 maps the vector part onto itself."""
434
+ jac = jax.jacobian(lambda v: Quaternion.from_scalar_vector(1.0, v).log().wxyz)(jnp.zeros(3))
435
+ expected = jnp.concatenate([jnp.zeros((1, 3)), jnp.eye(3)])
436
+ assert jnp.allclose(jac, expected)
@@ -79,6 +79,80 @@ def test_from_rotation_matrix_wrong_shape(do_jit):
79
79
  func(wrong_matrix)
80
80
 
81
81
 
82
+ # from_axis_angle
83
+ @pytest.mark.parametrize(
84
+ 'axis, angle, expected',
85
+ [
86
+ ([0.0, 0.0, 1.0], 0.0, [1.0, 0.0, 0.0, 0.0]),
87
+ ([1.0, 0.0, 0.0], jnp.pi, [0.0, 1.0, 0.0, 0.0]),
88
+ ([0.0, 2.0, 0.0], jnp.pi / 2, [jnp.sqrt(0.5), 0.0, jnp.sqrt(0.5), 0.0]),
89
+ ([0.0, 0.0, 1.0], -jnp.pi / 2, [jnp.sqrt(0.5), 0.0, 0.0, -jnp.sqrt(0.5)]),
90
+ ],
91
+ )
92
+ @pytest.mark.parametrize('do_jit', [False, True])
93
+ def test_from_axis_angle(axis, angle, expected, do_jit):
94
+ """Test from_axis_angle against known quaternions (axis need not be normalized)."""
95
+ func = Quaternion.from_axis_angle
96
+ if do_jit:
97
+ func = jax.jit(func)
98
+
99
+ q = func(jnp.array(axis), jnp.array(angle))
100
+ assert jnp.allclose(q.wxyz, jnp.array(expected), atol=1e-6)
101
+
102
+
103
+ @pytest.mark.parametrize('do_jit', [False, True])
104
+ def test_from_axis_angle_right_hand_rule(do_jit):
105
+ """A positive rotation about z maps x to y."""
106
+
107
+ def func(angle):
108
+ return Quaternion.from_axis_angle(jnp.array([0.0, 0.0, 1.0]), angle).rotate_vector(
109
+ jnp.array([1.0, 0.0, 0.0])
110
+ )
111
+
112
+ if do_jit:
113
+ func = jax.jit(func)
114
+
115
+ assert jnp.allclose(func(jnp.pi / 2), jnp.array([0.0, 1.0, 0.0]), atol=1e-6)
116
+
117
+
118
+ def test_from_axis_angle_consistency_with_matrix():
119
+ """The rotation matrix of from_axis_angle is Rodrigues' formula."""
120
+ axis = jnp.array([1.0, -2.0, 0.5])
121
+ angle = 0.7
122
+ n = axis / jnp.linalg.norm(axis)
123
+ k = jnp.array([[0, -n[2], n[1]], [n[2], 0, -n[0]], [-n[1], n[0], 0]])
124
+ expected = jnp.eye(3) + jnp.sin(angle) * k + (1 - jnp.cos(angle)) * k @ k
125
+
126
+ q = Quaternion.from_axis_angle(axis, angle)
127
+ assert jnp.allclose(q.to_rotation_matrix(), expected, atol=1e-6)
128
+
129
+
130
+ def test_from_axis_angle_broadcast():
131
+ """Axes and angles broadcast against each other."""
132
+ axes = jnp.eye(3) # (3, 3)
133
+ angles = jnp.array([[0.1], [0.2]]) # (2, 1)
134
+ q = Quaternion.from_axis_angle(axes, angles)
135
+ assert q.shape == (2, 3)
136
+ assert jnp.allclose(abs(q), 1.0, atol=1e-6)
137
+ expected = Quaternion.from_axis_angle(axes[2], angles[1, 0])
138
+ assert jnp.allclose(q[1, 2].wxyz, expected.wxyz, atol=1e-6)
139
+
140
+
141
+ def test_from_axis_angle_grad_at_zero():
142
+ """The derivative with respect to the angle is finite and exact at zero."""
143
+
144
+ def func(angle):
145
+ return Quaternion.from_axis_angle(jnp.array([0.0, 1.0, 0.0]), angle).wxyz
146
+
147
+ jac = jax.jacfwd(func)(0.0)
148
+ assert jnp.allclose(jac, jnp.array([0.0, 0.0, 0.5, 0.0]))
149
+
150
+
151
+ def test_from_axis_angle_wrong_shape():
152
+ with pytest.raises(ValueError, match='Axis must have shape'):
153
+ Quaternion.from_axis_angle(jnp.array([1.0, 0.0]), 0.1)
154
+
155
+
82
156
  # to_rotation_matrix
83
157
  @pytest.mark.parametrize('do_jit', [False, True])
84
158
  def test_to_rotation_matrix_identity(do_jit):
@@ -666,6 +666,7 @@ dev = [
666
666
  { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
667
667
  { name = "numpy", version = "2.4.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" },
668
668
  { name = "pillow" },
669
+ { name = "prek" },
669
670
  { name = "pytest" },
670
671
  { name = "pytest-notebook" },
671
672
  { name = "sphinx", version = "8.1.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
@@ -688,6 +689,9 @@ docs = [
688
689
  { name = "sphinx", version = "9.1.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" },
689
690
  { name = "sphinx-rtd-theme" },
690
691
  ]
692
+ lint = [
693
+ { name = "prek" },
694
+ ]
691
695
 
692
696
  [package.metadata]
693
697
  requires-dist = [
@@ -705,6 +709,7 @@ dev = [
705
709
  { name = "nbsphinx", specifier = ">=0.9.0" },
706
710
  { name = "numpy", specifier = ">=1.24.0" },
707
711
  { name = "pillow", specifier = ">=9.0.0" },
712
+ { name = "prek", specifier = ">=0.5.4" },
708
713
  { name = "pytest", specifier = ">=9.0" },
709
714
  { name = "pytest-notebook", specifier = ">=0.10" },
710
715
  { name = "sphinx", specifier = ">=7.0.0" },
@@ -721,6 +726,7 @@ docs = [
721
726
  { name = "sphinx", specifier = ">=7.0.0" },
722
727
  { name = "sphinx-rtd-theme", specifier = ">=1.3.0" },
723
728
  ]
729
+ lint = [{ name = "prek", specifier = ">=0.5.4" }]
724
730
 
725
731
  [[package]]
726
732
  name = "fonttools"
@@ -2778,6 +2784,22 @@ wheels = [
2778
2784
  { url = "https://files.pythonhosted.org/packages/54/20/4d324d65cc6d9205fabedc306948156824eb9f0ee1633355a8f7ec5c66bf/pluggy-1.6.0-py3-none-any.whl", hash = "sha256:e920276dd6813095e9377c0bc5566d94c932c33b27a3e3945d8389c374dd4746", size = 20538, upload-time = "2025-05-15T12:30:06.134Z" },
2779
2785
  ]
2780
2786
 
2787
+ [[package]]
2788
+ name = "prek"
2789
+ version = "0.5.4"
2790
+ source = { registry = "https://pypi.org/simple" }
2791
+ sdist = { url = "https://files.pythonhosted.org/packages/30/f1/8f3143530bf43fa82f8873f4e536f318b6143dd5c09171aba0d465e21f04/prek-0.5.4.tar.gz", hash = "sha256:7f5bc061880141e362ef422e52188d0090794c0622e442e3157686d1e342780a", size = 565182, upload-time = "2026-09-28T04:54:43.453Z" }
2792
+ wheels = [
2793
+ { url = "https://files.pythonhosted.org/packages/78/a7/f3a4a252c0846c5846dd86118ee0c1fddb7381dc35e3d44c2a5f3cea1757/prek-0.5.4-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:550ecaae0601f9405526c866a57cc6882ed1e9d8696ac1117b8703bb37d087eb", size = 6011113, upload-time = "2026-09-28T04:54:28.107Z" },
2794
+ { url = "https://files.pythonhosted.org/packages/1b/e3/26638cf090b620e63e60a639b8e7d7eb38b9fbea74a0827a182fc922cf8e/prek-0.5.4-py3-none-macosx_11_0_arm64.whl", hash = "sha256:a7ef67cc23e23c6e9c6ca0b27a7b12c19f16774d36b1c88ed035f2639995718d", size = 5534338, upload-time = "2026-09-28T04:54:30.214Z" },
2795
+ { url = "https://files.pythonhosted.org/packages/a0/08/64b668eb088204650b96183cea739b56e6b179130c4398a8a83cbe9b125a/prek-0.5.4-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.musllinux_1_1_aarch64.whl", hash = "sha256:59f2206850f643bafe46b28b4431525f8dcc5b7a763f271179bff10b8157870d", size = 5840937, upload-time = "2026-09-28T04:54:32.226Z" },
2796
+ { url = "https://files.pythonhosted.org/packages/1f/ee/a0f847194eb6a2b67921b4c94c4fbf3f33a785affb190c6a2097344dbdd6/prek-0.5.4-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:7d5b0cab1219a7bc9cc90785c7436f563765440616543ede7dfa3ba5c332cede", size = 6234354, upload-time = "2026-09-28T04:54:34.176Z" },
2797
+ { url = "https://files.pythonhosted.org/packages/32/9a/14186d90ed29b46650452068d6bd1ac233ca0ea41a3b35081c37f6b77b1c/prek-0.5.4-py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:8fc22db8132dab9aed4e152bf6a58a7f791f4c91c74305d65d96c48b8ec0bc62", size = 5845649, upload-time = "2026-09-28T04:54:36.252Z" },
2798
+ { url = "https://files.pythonhosted.org/packages/e0/88/ff5bbbada4098f59740ab767f5c2f82d4110b90ea7db309a9bcf21f3dcb9/prek-0.5.4-py3-none-musllinux_1_1_x86_64.whl", hash = "sha256:19d4f3986dcd1dfa15f9417e32a75fb2ba084156dbd7ace161f0764281b8064b", size = 6351516, upload-time = "2026-09-28T04:54:38.539Z" },
2799
+ { url = "https://files.pythonhosted.org/packages/c8/1a/0cc410232ef42750b3bf34babd6def4a3943ec98259d9fecda3b26fcd63f/prek-0.5.4-py3-none-win_amd64.whl", hash = "sha256:3100233ba802eff6a4b09bc045c8f50d599bb9d810fa234823bc1642938b224f", size = 5755956, upload-time = "2026-09-28T04:54:40.227Z" },
2800
+ { url = "https://files.pythonhosted.org/packages/98/78/3fd53cc54468db37462885d3139eae40bb469d51beae4575c986b190ab6e/prek-0.5.4-py3-none-win_arm64.whl", hash = "sha256:9fc4ad2d647f70a8a3a9900849e1dd89472cf5a71d830f083b62cf87a0406e27", size = 5502478, upload-time = "2026-09-28T04:54:41.973Z" },
2801
+ ]
2802
+
2781
2803
  [[package]]
2782
2804
  name = "prometheus-client"
2783
2805
  version = "0.24.1"
@@ -1,67 +0,0 @@
1
- name: CI
2
-
3
- on:
4
- push:
5
- branches:
6
- - main
7
- pull_request:
8
- branches:
9
- - main
10
-
11
- concurrency:
12
- group: ${{ github.workflow }}-${{ github.ref }}
13
- cancel-in-progress: true
14
-
15
- jobs:
16
-
17
- pre-commit:
18
- name: Pre-commit
19
- runs-on: ubuntu-latest
20
- steps:
21
- - uses: actions/checkout@v4
22
- - name: Set up Python
23
- uses: actions/setup-python@v5
24
- with:
25
- python-version: '3.x'
26
- - name: Install dependencies
27
- run: |
28
- python -m pip install --upgrade pip
29
- python -m pip install pre-commit
30
- - name: Run pre-commit
31
- run: pre-commit run --all-files --show-diff-on-failure --color=always
32
-
33
- test:
34
- name: Run tests on Python ${{ matrix.python-version }}
35
- runs-on: ubuntu-latest
36
- strategy:
37
- matrix:
38
- python-version: ['3.10', '3.11', '3.12', '3.13', '3.14']
39
-
40
- steps:
41
- - name: Checkout code
42
- uses: actions/checkout@v4
43
- with:
44
- fetch-depth: 0
45
-
46
- - name: Set up Python ${{ matrix.python-version }}
47
- uses: actions/setup-python@v5
48
- with:
49
- python-version: ${{ matrix.python-version }}
50
-
51
- - name: Install uv
52
- uses: astral-sh/setup-uv@v6
53
- with:
54
- enable-cache: true
55
-
56
- - name: Build wheel
57
- run: uv build --wheel
58
-
59
- - name: Install project's wheel and dependencies
60
- run: |
61
- uv sync --locked --no-install-project
62
- uv pip install dist/fastquat-*.whl
63
-
64
- - name: Run tests
65
- run: |
66
- source .venv/bin/activate
67
- pytest -v
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes