复制安装命令
用 Codex 或 Claude 安装复制这段 Prompt,粘贴到 Codex、Claude 或其他助手里,让它先审查 Skill 页面再帮你安装。
复制前请先查看来源、License 和安全提示。
Unitful Quantities in JAX
用 Codex 或 Claude 安装复制这段 Prompt,粘贴到 Codex、Claude 或其他助手里,让它先审查 Skill 页面再帮你安装。
复制前请先查看来源、License 和安全提示。
来源文件:README.md
Unxt is unitful quantities and calculations in JAX, built on Equinox and Quax.
Unxt supports JAX's compelling features:
jit)vmap, etc.)grad, jacobian, hessian)And best of all, unxt doesn't force you to use special unit-compatible re-exports of JAX libraries. You can use unxt with existing JAX code, and with quax's simple decorator, JAX will work with unxt.Quantity.
pip install unxt
uvuv add unxt
pip install git+https://github.com/GalacticDynamics/unxt.git
cd /path/to/parent
git clone https://github.com/GalacticDynamics/unxt.git
cd unxt
pip install -e . # editable mode
For full documentation, including installation instructions, tutorials, and API reference, please see the unxt docs. This README provides a brief overview and some quick examples.
Dimensions represent the physical type of a quantity, such as length, time, or mass.
>>> import unxt as u
Create dimensions from strings:
>>> u.dimension("length")
PhysicalType('length')
Dimensions support mathematical expressions:
>>> u.dimension("length / time")
PhysicalType({'speed', 'velocity'})
Multi-word dimension names require parentheses in expressions:
>>> u.dimension("(amount of substance) / (time)")
PhysicalType('catalytic activity')
Units specify the scale and dimension of measurements.
>>> meter = u.unit("m")
>>> meter
Unit("m")
Units can be combined, either inside the expression string or by arithmetic on Unit objects:
>>> u.unit("km/h") # in the expression
Unit("km / h")
>>> u.unit("km") / u.unit("h") # via arithmetic
Unit("km / h")
Get the dimension of a unit:
>>> u.dimension_of(meter)
PhysicalType('length')
Unit systems define consistent sets of base units for specific domains. unxt provides built-in unit systems and tools for creating custom ones.
>>> u.unitsystem("si") # SI (International System of Units)
unitsystem(m, kg, s, mol, A, K, cd, rad)
>>> u.unitsystem("cgs") # CGS (centimeter-gram-second)
unitsystem(cm, g, s, dyn, erg, Ba, P, St, rad)
>>> u.unitsystem("galactic") # galactic (astrophysics)
unitsystem(kpc, Myr, solMass, rad)
Once you have a unit system, you can get units for any physical dimension by indexing the system:
>>> usys = u.unitsystem("si")
>>> usys["length"]
Unit("m")
Create custom unit systems by specifying base units:
>>> custom_usys = u.unitsystem("km", "h", "tonne", "degree")
>>> custom_usys
unitsystem(km, h, t, deg)
Derived units are then available by dimension:
>>> custom_usys["velocity"]
Unit("km / h")
For domains like gravitational dynamics, use dynamical unit systems where $G = 1$. Specify only 2 of (length, time, mass); the third is computed to make $G = 1$.
>>> from unxt.unitsystems import DynamicalSimUSysFlag
>>> dyn_usys = u.unitsystem(DynamicalSimUSysFlag, "kpc", "Myr")
>>> dyn_usys
LengthMassTimeUnitSystem(length=Unit("kpc"),
mass=Unit("1.49828e+10 kpc3 s2 kg / (Myr2 m3)"), time=Unit("Myr"))
The mass unit is the derived one — an exact composite expression, not a rounded label:
>>> dyn_usys["mass"]
Unit("1.49828e+10 kpc3 s2 kg / (Myr2 m3)")
Quantities combine values with units, providing type-safe unitful arithmetic.
Quantity (u.Q) is the lightweight, non-parametric default: a single class — and a single JAX pytree type — for all physical dimensions. ParametricQuantity (up.PQ) adds runtime dimension checking and dimension-specific plum dispatch by encoding each dimension in its own on-the-fly class (and pytree type), which grows the type/dispatch surface and adds per-construction overhead. (This is not about jax.jit cache misses: the unit is static, so a jitted function specializes per unit with either class — that part is inherent.) See the Quantity guide for full details; upgrading from an earlier version? See the migration guide.
>>> import jax.numpy as jnp
>>> x = u.Q(jnp.arange(1, 5, dtype=float), "km")
>>> x
Quantity(Array([1., 2., 3., 4.], dtype=float32...), unit='km')
The constituent value and unit are accessible as attributes:
>>> x.value
Array([1., 2., 3., 4.], dtype=float32...)
>>> x.unit
Unit("km")
Quantity objects obey the rules of unitful arithmetic — addition, subtraction, multiplication, division, and exponentiation:
>>> x + x
Quantity(Array([2., 4., 6., 8.], dtype=float32...), unit='km')
>>> 2 * x
Quantity(Array([2., 4., 6., 8.], dtype=float32...), unit='km')
>>> y = u.Q(jnp.arange(4, 8, dtype=float), "yr")
>>> x / y
Quantity(
Array([0.25 , 0.4 , 0.5 , 0.5714286], dtype=float32...), unit='km / yr'
)
>>> x**2
Quantity(Array([ 1., 4., 9., 16.], dtype=float32...), unit='km2')
Operations are unit-checked, so mixing incompatible dimensions raises:
>>> try:
... x + y
... except Exception as e:
... print(e)
'yr' (time) and 'km' (length) are not convertible
Quantities can be converted to different units, by function or by method:
>>> u.uconvert("m", x) # via function
Quantity(Array([1000., 2000., 3000., 4000.], dtype=float32...), unit='m')
>>> x.uconvert("m") # via method
Quantity(Array([1000., 2000., 3000., 4000.], dtype=float32...), unit='m')
ParametricQuantity — from the separate unxts.parametric package (pip install unxts.parametric, imported below as up) — adds runtime dimension checking on construction. Use up.PQ["length"] to create a parametric type that raises if the unit's physical type does not match:
>>> import unxts.parametric as up
>>> LengthQuantity = up.PQ["length"]
>>> LengthQuantity(2, "km")
ParametricQuantity(Array(2, dtype=int32...), unit='km')
A unit whose physical type does not match the parameter is rejected:
>>> try:
... LengthQuantity(2, "s")
... except ValueError as e:
... print(e)
Physical type mismatch.
By contrast, the default u.Q["length"] accepts the subscript but does not check dimensions — it silently builds a Quantity with the mismatched unit:
>>> u.Q["length"](2, "s")
Quantity(Array(2, dtype=int32...), unit='s')
Use up.PQ["length"] when you need the runtime guard. See the unxts.parametric guide for the full API.
Quantity (aliased as u.Q) is the default lightweight class. It does not do runtime dimension checking on construction, which makes it the fastest option for performance-critical code:
>>> bq = u.quantity.Quantity(jnp.array([1.0, 2.0, 3.0]), "m")
>>> bq
Quantity(Array([1., 2., 3.], dtype=float32...), unit='m')
>>> bq * 2
Quantity(Array([2., 4., 6.], dtype=float32...), unit='m')
Angle is a specialized quantity with wrapping support for angular values:
>>> theta = u.Angle(jnp.array([0, 90, 180, 270, 360]), "deg")
>>> theta
Angle(Array([ 0, 90, 180, 270, 360], dtype=int32...), unit='deg')
Angles can optionally be wrapped into a specified range:
>>> angle = u.Angle(jnp.array([370, -10]), "deg")
>>> angle.wrap_to(u.Q(0, "deg"), u.Q(360, "deg"))
Angle(Array([ 10, 350], dtype=int32...), unit='deg')
For static configuration values (e.g., JAX static arguments), use StaticQuantity, which stores NumPy values and rejects JAX arrays:
>>> import numpy as np
>>> from functools import partial
>>> import jax
>>> cfg = u.StaticQuantity(np.array([1.0, 2.0]), "m")
>>> @partial(jax.jit, static_argnames=("q",))
... def add(x, q):
... return x + jnp.asarray(q.value)
>>> add(1.0, cfg)
Array([2., 3.], dtype=float32...)
If you want a Quantity that keeps a static value but still participates in regular arithmetic, wrap the value with StaticValue. Arithmetic behaves like the wrapped array, and StaticValue + StaticValue returns a StaticValue. Equality between two StaticValues (== / !=) returns a scalar bool — which is what makes a StaticValue-backed quantity hashable and usable as a jax.jit static argument. Ordering (<, <=, >, >=), and == / != against a raw array, return element-wise NumPy boolean arrays:
>>> sv = u.quantity.StaticValue(np.array([1.0, 2.0]))
>>> q_static = u.Q(sv, "m")
>>> q = u.Q(jnp.array([3.0, 4.0]), "m")
>>> q_static + q
Quantity(Array([4., 6.], dtype=float32...), unit='m')
Equality between two StaticValues is a scalar bool:
>>> sv2 = u.quantity.StaticValue(np.array([2.0, 1.0]))
>>> sv == sv2
False
>>> sv == u.quantity.StaticValue(np.array([1.0, 2.0]))
True
Ordering, and equality against a raw array, are element-wise NumPy boolean arrays:
>>> sv < sv2
array([ True, False])
>>> sv == np.array([1.0, 2.0])
array([ True, True])
unxt is built on quax, which enables custom array-ish objects in JAX. For convenience we use the quaxed library, which is just a quax.quaxify wrapper around jax to avoid boilerplate code.
[!NOTE]
Using
quaxedis optional. You can directly usequaxify, and even apply it to the top-level function instead of individual functions.
Using the x quantity from the earlier examples:
>>> from quaxed import grad, vmap
>>> import quaxed.numpy as qnp
>>> qnp.square(x)
Quantity(Array([ 1., 4., 9., 16.], dtype=float32...), unit='km2')
>>> qnp.power(x, 3)
Quantity(Array([ 1., 8., 27., 64.], dtype=float32...), unit='km3')
>>> vmap(grad(lambda x: x**3))(x)
Quantity(Array([ 3., 12., 27., 48.], dtype=float32...), unit='km2')
See the documentation for more examples and details of JIT and AD
If you found this library to be useful and want to support the development and maintenance of lower-level code libraries for the scientific community, please consider citing this work.
We welcome contributions! Contributions are how open source projects improve and grow.
To contribute to unxt, please fork the repository, make a development branch, develop on that branch, then open a pull request from the branch in your fork to main.
To report bugs, request features, or suggest other ideas, please open an issue.
For more information, see CONTRIBUTING.md.
name: unxt
description: >
Use when writing, reviewing, or debugging code that imports unxt (Quantity, ParametricQuantity, units, dims, unit systems), or that passes a unxt Quantity through JAX/quax/quaxed functions. Also use when a dimension mismatch error like "'yr' (time) and 'km' (length) are not convertible" appears, when plum raises an ambiguous-dispatch error on unit()/dimension()/ convert(), when Quantity["length"] silently accepts a mismatched unit instead of raising, when code reaches for a Quantity's private `_mk` constructor, when a StaticQuantity comparison behaves unit-blind, or when upgrading code that still references the deprecated BareQuantity.unxt gives JAX unitful quantities: Quantity is a quax.ArrayValue (an Equinox PyTree), so it flows through jax.jit/vmap/grad like any other JAX array-ish type, while carrying a unit and enforcing unit-safe arithmetic.
Checked against unxt 2.0.x (quax>=0.4.2, quax-blocks>=0.5.0, quaxed>=0.10.5, plum-dispatch>=2.7.0, astropy>=7.1), Python >=3.12. Docs: https://unxt.readthedocs.io/en/.
Read the quax, quaxed, and quax-blocks skills first for anything about quaxify, dispatch resolution, or the mixin operator overloads — this skill covers only what's specific to unxt, and doesn't restate any of it.
>>> import jax.numpy as jnp
>>> import unxt as u
>>> velocity = u.Q(30.0, "m/s")
>>> time = u.Q(2.0, "s")
>>> velocity * time
Quantity(Array(60., dtype=float32, ...), unit='m')
>>> u.uconvert("km", velocity * time)
Quantity(Array(0.06, dtype=float32, ...), unit='km')
Quantity vs ParametricQuantity vs StaticQuantity vs (deprecated) BareQuantity| Class | Package | Dimension in type? | Checked at construction? | Use when |
|---|---|---|---|---|
Quantity/u.Q | unxt (default) | no — one pytree type for every dimension | no | almost always; the fast, general default |
ParametricQuantity/up.PQ | unxts.parametric (opt-in) | yes, e.g. PQ["length"] | yes, raises on mismatch | you want a runtime guard or dimension-specific dispatch, and can accept a distinct pytree type per dimension |
StaticQuantity | unxt | no | — | value must be jax.jit(static_argnames=...)-hashable; equality is unit-label-based, not physical-equivalence |
BareQuantity | unxt | — | — | deprecated, alias of Quantity — don't use in new code, see the migration guide |
The subscript-without-checking trap: u.Q["length"] accepts the subscript syntax but performs no dimension check — it silently builds a Quantity with whatever unit you give it:
>>> u.Q["length"](2, "s") # wrong dimension, no error
Quantity(Array(2, dtype=int32, ...), unit='s')
>>> import unxts.parametric as up
>>> up.PQ["length"](2, "s") # doctest: +SKIP
# ValueError: Physical type mismatch.
If you need the guard, you need unxts.parametric's PQ, not u.Q[...].
The functional API is primary (the OO methods just call it). Argument order is inspired by Unitful.jl: operator first, operand last — read uconvert("cm", q) as "convertto cm":
>>> q = u.Q(1, "m")
>>> u.uconvert("cm", q) # function form — operator first
Quantity(Array(100., dtype=float32, ...), unit='cm')
>>> q.uconvert("cm") # OO form — same result
Quantity(Array(100., dtype=float32, ...), unit='cm')
Don't write uconvert(q, "cm") — that's the wrong order and, depending on dispatch coverage, may raise a plum ambiguity/no-method error instead of silently doing the wrong thing.
_mk is private API — do not use unless you mean itQuantity._mk (and type-specific overrides like QuantityMatrix._mk) is not exported and not covered by semver. It writes the value/unit fields directly and skips both the plum-dispatched value/unit converters and __check_init__ — the checks that make normal construction safe. It exists purely as a ~50x-faster hot-path constructor for code that has already proven its inputs are normalised.
Warn the user explicitly before introducing _mk in code written for them. It is only safe when:
value is already the right array type/dtype for this quantity's storage, andunit is already a real AbstractUnit instance (not a string), andGet any of that wrong and you get a Quantity that looks valid but violates its own invariants — e.g. a unit that's still a string, silently breaking every downstream dispatch that expects AbstractUnit. StaticQuantity overrides _mk back to the checked constructor for exactly this reason: its converter is load-bearing, not redundant. If a value/unit pair isn't provably pre-normalised, use the normal constructor or revalue, not _mk.
+/-, not just anything unexpectedu.dimension(...) parses a small expression grammar: * / ** () work, but unary +/- raise on purpose ("dimensions are invariant under negation") — this is a deliberate rejection, not a missing feature to "fix":
>>> u.dimension("length / time")
PhysicalType({'speed', 'velocity'})
StaticQuantity equality is unit-label-based, not physical. == compares unit labels (same_unit_label), not physical equivalence — two quantities with the same value but different unit spelling for the same physical unit can compare unequal, because equality must stay a valid jax.jit static_argnames key. Use unxt.equivalent/is_equivalent for physical-equivalence comparison instead.u.Q["length"] doesn't check dimensions. See above — that's unxts.parametric.PQ's job, not Quantity's.si, cgs, dimensionless) are deliberately shared, immutable objects — code that used to be able to corrupt them by mutation was a bug (fixed in #704/#718); don't reintroduce mutable state on these.| Need | Package |
|---|---|
Runtime-checked, dimension-typed quantities (PQ["length"]) | unxts.parametric |
Heterogeneous-unit matrices/vectors (QuantityMatrix/QM, UnitsMatrix) | unxts.linalg |
gala.units.UnitSystem interop | unxts.interop.gala |
Plotting Quantity with matplotlib | unxts.interop.matplotlib |
| xarray accessors for quantities | unxts.interop.xarray |
| Hypothesis strategies for property-based tests | unxts.hypothesis |
| Minimal-dependency abstract dispatch API only | unxts.api |
Each of these will get its own skills/<pkg>/SKILL.md in time; until then, their docs/packages/<pkg>/ pages are the reference.
| Symptom | Cause / fix |
|---|---|
'yr' (time) and 'km' (length) are not convertible | Arithmetic/uconvert between incompatible dimensions — this is unxt working correctly; convert one side first or check you meant the operation. |
Physical type mismatch. from ParametricQuantity/PQ[...] construction | The unit you passed doesn't match the type parameter's dimension. Only PQ checks this — Quantity/u.Q silently accepts it. |
Ambiguous-dispatch / no-method error from unit(), dimension(), or convert | You passed a type with no registered conversion. Check <func>.methods to see what's registered, or convert to a supported type first (str, AbstractUnit, astropy.units.Unit, ...). |
u.Q["length"](2, "s") returns a Quantity with the wrong dimension, no error | Expected — Quantity doesn't check dimensions on subscript construction. Use unxts.parametric.PQ["length"] if you need the guard. |
A Quantity built via _mk behaves wrong downstream (wrong dtype, unit still a string, dispatch misses) | _mk was called with unnormalised input. Don't use _mk outside code that has already proven normalisation; use the checked constructor or revalue. |
A jax.jit/quaxify outer-wrapper function is far slower than expected, despite following the perf guide | It's likely being rebuilt per call (inside a loop, or a fresh @jax.jit def ... per function call) instead of built once. jax.jit caches on function identity, not argument equality — a fresh closure is a compile-cache miss every time. |
Doctest in a docstring/README/docs/*.md fails on dtype (dtype=float32 vs dtype=float) | Sybil matches output exactly. Match the real dtype JAX produces, don't approximate it. |
| A warning you didn't expect fails the test suite | filterwarnings = ["error", ...] in pyproject.toml — either fix the cause or add a scoped, justified ignore; don't silence broadly. |
unxt v2.0 restructured the quantity hierarchy: BareQuantity is deprecated in favor of plain Quantity; dimension-parametrized quantities moved out to the separate unxts.parametric package (PQ); several previously-unxt-*-hyphenated packages have canonical unxts.* replacements (the hyphenated packages are now back-compat shims). See the migration guide for the full v1→v2 mapping before assuming an older code sample is still current.
评论 (0)
暂无评论,成为第一个评论者吧!