eft4py 0.0.2__tar.gz → 0.0.4__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.
- {eft4py-0.0.2 → eft4py-0.0.4}/PKG-INFO +8 -1
- {eft4py-0.0.2 → eft4py-0.0.4}/pyproject.toml +9 -1
- {eft4py-0.0.2 → eft4py-0.0.4}/pyproject.toml.orig +9 -1
- eft4py-0.0.4/src/eft4py/__init__.py +94 -0
- eft4py-0.0.4/src/eft4py/backend/__init__.py +138 -0
- eft4py-0.0.4/src/eft4py/core/__init__.py +30 -0
- eft4py-0.0.4/src/eft4py/core/protocols.py +167 -0
- eft4py-0.0.4/src/eft4py/core/types.py +175 -0
- eft4py-0.0.4/src/eft4py/models/__init__.py +5 -0
- eft4py-0.0.4/src/eft4py/models/horndeski.py +513 -0
- eft4py-0.0.4/src/eft4py/physics/__init__.py +44 -0
- eft4py-0.0.4/src/eft4py/physics/alpha_functions.py +206 -0
- eft4py-0.0.4/src/eft4py/physics/background.py +160 -0
- eft4py-0.0.4/src/eft4py/physics/compiled.py +288 -0
- eft4py-0.0.4/src/eft4py/physics/cosmology.py +173 -0
- eft4py-0.0.4/src/eft4py/physics/dark_energy.py +154 -0
- eft4py-0.0.4/src/eft4py/physics/derivatives.py +185 -0
- eft4py-0.0.4/src/eft4py/physics/errors.py +31 -0
- eft4py-0.0.4/src/eft4py/physics/evolution_coeffs.py +129 -0
- eft4py-0.0.4/src/eft4py/physics/initial_conditions.py +152 -0
- eft4py-0.0.4/src/eft4py/physics/neutrinos.py +105 -0
- eft4py-0.0.4/src/eft4py/physics/quasi_static.py +134 -0
- eft4py-0.0.4/src/eft4py/physics/sound_horizon.py +325 -0
- eft4py-0.0.4/src/eft4py/physics/stability.py +469 -0
- eft4py-0.0.4/src/eft4py/physics/symbolic.py +182 -0
- eft4py-0.0.4/src/eft4py/solvers/__init__.py +65 -0
- eft4py-0.0.4/src/eft4py/solvers/background.py +701 -0
- eft4py-0.0.4/src/eft4py/solvers/linear/__init__.py +70 -0
- eft4py-0.0.4/src/eft4py/solvers/linear/emulator.py +361 -0
- eft4py-0.0.4/src/eft4py/solvers/linear/growth.py +275 -0
- eft4py-0.0.4/src/eft4py/solvers/linear/limber.py +362 -0
- eft4py-0.0.4/src/eft4py/solvers/linear/observables.py +157 -0
- eft4py-0.0.4/src/eft4py/solvers/linear/potentials.py +178 -0
- eft4py-0.0.4/src/eft4py/solvers/linear/power.py +73 -0
- eft4py-0.0.4/src/eft4py/solvers/ode_systems/__init__.py +10 -0
- eft4py-0.0.4/src/eft4py/solvers/ode_systems/base.py +212 -0
- eft4py-0.0.4/src/eft4py/solvers/ode_systems/horndeski.py +591 -0
- eft4py-0.0.4/src/eft4py/solvers/solution.py +527 -0
- eft4py-0.0.4/src/eft4py/solvers/strategies/__init__.py +31 -0
- eft4py-0.0.4/src/eft4py/solvers/strategies/base.py +130 -0
- eft4py-0.0.4/src/eft4py/solvers/strategies/jax.py +295 -0
- eft4py-0.0.4/src/eft4py/solvers/strategies/scipy.py +220 -0
- eft4py-0.0.4/src/eft4py/utils/__init__.py +9 -0
- eft4py-0.0.4/src/eft4py/utils/_classy.py +190 -0
- eft4py-0.0.4/src/eft4py/utils/hermite.py +99 -0
- eft4py-0.0.4/src/eft4py/utils/jax_config.py +33 -0
- eft4py-0.0.4/src/eft4py/utils/nowiggle.py +252 -0
- eft4py-0.0.4/src/eft4py/utils/symbols.py +201 -0
- {eft4py-0.0.2 → eft4py-0.0.4}/README.md +0 -0
- /eft4py-0.0.2/src/eft4py/__init__.py → /eft4py-0.0.4/src/eft4py/py.typed +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: eft4py
|
|
3
|
-
Version: 0.0.
|
|
3
|
+
Version: 0.0.4
|
|
4
4
|
Summary: An Effective Field Theory toolbox for Python
|
|
5
5
|
Author: Rodrigo Calderon
|
|
6
6
|
Author-email: Rodrigo Calderon <calderon.cosmology@gmail.com>
|
|
@@ -11,12 +11,19 @@ Requires-Dist: pytest>=7.4 ; extra == 'dev'
|
|
|
11
11
|
Requires-Dist: pytest-cov>=4.1 ; extra == 'dev'
|
|
12
12
|
Requires-Dist: mypy>=1.10 ; extra == 'dev'
|
|
13
13
|
Requires-Dist: ruff>=0.5 ; extra == 'dev'
|
|
14
|
+
Requires-Dist: mkdocs-material>=9.5 ; extra == 'docs'
|
|
15
|
+
Requires-Dist: mkdocstrings[python]>=0.26 ; extra == 'docs'
|
|
16
|
+
Requires-Dist: mkdocs-marimo>=0.2.1 ; extra == 'docs'
|
|
17
|
+
Requires-Dist: mkdocs-jupyter>=0.25 ; extra == 'docs'
|
|
14
18
|
Requires-Dist: jax>=0.4.0 ; extra == 'jax'
|
|
15
19
|
Requires-Dist: jaxlib>=0.4.0 ; extra == 'jax'
|
|
20
|
+
Requires-Dist: diffrax>=0.5.0 ; extra == 'jax'
|
|
21
|
+
Requires-Dist: optimistix>=0.0.10 ; extra == 'jax'
|
|
16
22
|
Requires-Dist: matplotlib>=3.7 ; extra == 'notebooks'
|
|
17
23
|
Requires-Dist: jupyter>=1.0 ; extra == 'notebooks'
|
|
18
24
|
Requires-Python: >=3.12
|
|
19
25
|
Provides-Extra: dev
|
|
26
|
+
Provides-Extra: docs
|
|
20
27
|
Provides-Extra: jax
|
|
21
28
|
Provides-Extra: notebooks
|
|
22
29
|
Description-Content-Type: text/markdown
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "eft4py"
|
|
3
|
-
version = "0.0.
|
|
3
|
+
version = "0.0.4"
|
|
4
4
|
description = "An Effective Field Theory toolbox for Python"
|
|
5
5
|
readme = "README.md"
|
|
6
6
|
requires-python = ">=3.12"
|
|
@@ -28,6 +28,14 @@ notebooks = [
|
|
|
28
28
|
jax = [
|
|
29
29
|
"jax>=0.4.0",
|
|
30
30
|
"jaxlib>=0.4.0",
|
|
31
|
+
"diffrax>=0.5.0",
|
|
32
|
+
"optimistix>=0.0.10",
|
|
33
|
+
]
|
|
34
|
+
docs = [
|
|
35
|
+
"mkdocs-material>=9.5",
|
|
36
|
+
"mkdocstrings[python]>=0.26",
|
|
37
|
+
"mkdocs-marimo>=0.2.1",
|
|
38
|
+
"mkdocs-jupyter>=0.25",
|
|
31
39
|
]
|
|
32
40
|
|
|
33
41
|
[project.scripts]
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "eft4py"
|
|
3
|
-
version = "0.0.
|
|
3
|
+
version = "0.0.4"
|
|
4
4
|
description = "An Effective Field Theory toolbox for Python"
|
|
5
5
|
readme = "README.md"
|
|
6
6
|
authors = [
|
|
@@ -27,6 +27,14 @@ notebooks = [
|
|
|
27
27
|
jax = [
|
|
28
28
|
"jax>=0.4.0",
|
|
29
29
|
"jaxlib>=0.4.0",
|
|
30
|
+
"diffrax>=0.5.0",
|
|
31
|
+
"optimistix>=0.0.10",
|
|
32
|
+
]
|
|
33
|
+
docs = [
|
|
34
|
+
"mkdocs-material>=9.5",
|
|
35
|
+
"mkdocstrings[python]>=0.26",
|
|
36
|
+
"mkdocs-marimo>=0.2.1",
|
|
37
|
+
"mkdocs-jupyter>=0.25",
|
|
30
38
|
]
|
|
31
39
|
|
|
32
40
|
[project.scripts]
|
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
r"""eft4py: Horndeski dark-energy backgrounds in Python, validated against hi_class.
|
|
2
|
+
|
|
3
|
+
Define a scalar-tensor theory through its Lagrangian functions
|
|
4
|
+
$G_2, G_3, G_4, G_5$ of $\phi$ and $X$, then evolve the cosmological
|
|
5
|
+
background in conformal time $\tau$ following the hi_class conventions.
|
|
6
|
+
|
|
7
|
+
The essentials are available at the top level:
|
|
8
|
+
|
|
9
|
+
- [`HorndeskiModel`][eft4py.HorndeskiModel]: the theory.
|
|
10
|
+
- [`Cosmology`][eft4py.Cosmology]: the fluid content ($h$, $\Omega_m$,
|
|
11
|
+
$\Omega_r$) the scalar field evolves in.
|
|
12
|
+
- [`ScipyBackground`][eft4py.ScipyBackground] and
|
|
13
|
+
[`JaxBackground`][eft4py.JaxBackground]: background solutions from a model
|
|
14
|
+
and a cosmology (initial conditions, a solve on a grid in $\ln a$, and
|
|
15
|
+
shooting a coupling so that $H(a=1) = H_0$), with
|
|
16
|
+
[`frozen_field`][eft4py.frozen_field] for a field initially at rest.
|
|
17
|
+
`JaxBackground` needs the `jax` extra, imported only when it is created.
|
|
18
|
+
- [`HorndeskiODE`][eft4py.HorndeskiODE]: the background equations, for
|
|
19
|
+
lower-level control.
|
|
20
|
+
- [`HiClassState`][eft4py.HiClassState]: a background state (initial
|
|
21
|
+
conditions and solutions).
|
|
22
|
+
- [`ScipySolver`][eft4py.ScipySolver] and, with the `jax` extra,
|
|
23
|
+
`JAXSolver`: integrators.
|
|
24
|
+
- [`get_symbols`][eft4py.get_symbols] and
|
|
25
|
+
[`SymbolTable`][eft4py.SymbolTable]: the symbols $\phi$, $X$,
|
|
26
|
+
$M_{\rm pl}$ used to write models.
|
|
27
|
+
- [`SingularSystemError`][eft4py.SingularSystemError]: raised when the
|
|
28
|
+
scalar's kinetic structure degenerates.
|
|
29
|
+
|
|
30
|
+
Advanced functionality lives in the subpackages: `eft4py.physics` (symbolic
|
|
31
|
+
and compiled coefficients, dark energy, $\alpha$-functions), `eft4py.models`,
|
|
32
|
+
`eft4py.solvers`, `eft4py.backend` and `eft4py.utils`.
|
|
33
|
+
|
|
34
|
+
Examples:
|
|
35
|
+
>>> from eft4py import HorndeskiModel, get_symbols
|
|
36
|
+
>>> s = get_symbols()
|
|
37
|
+
>>> model = HorndeskiModel(G2=s.X - 0.5 * s.phi**2, G4=s.Mpl**2 / 2, symbols=s)
|
|
38
|
+
"""
|
|
39
|
+
|
|
40
|
+
from importlib import metadata as _metadata
|
|
41
|
+
|
|
42
|
+
from .core.types import HiClassState
|
|
43
|
+
from .models.horndeski import HorndeskiModel
|
|
44
|
+
from .physics.cosmology import Cosmology
|
|
45
|
+
from .physics.errors import SingularSystemError
|
|
46
|
+
from .solvers.background import JaxBackground, ScipyBackground, frozen_field
|
|
47
|
+
from .solvers.ode_systems.horndeski import HorndeskiODE
|
|
48
|
+
from .solvers.strategies.scipy import ScipySolver
|
|
49
|
+
from .utils.symbols import SymbolTable, get_symbols
|
|
50
|
+
|
|
51
|
+
try:
|
|
52
|
+
__version__ = _metadata.version("eft4py")
|
|
53
|
+
except _metadata.PackageNotFoundError: # pragma: no cover - uninstalled source tree
|
|
54
|
+
__version__ = "0.0.0+unknown"
|
|
55
|
+
|
|
56
|
+
DOCS_URL = "https://rcalderonb6.github.io/eft4py/"
|
|
57
|
+
"""Address of the documentation website."""
|
|
58
|
+
|
|
59
|
+
__all__ = [
|
|
60
|
+
"__version__",
|
|
61
|
+
"HorndeskiModel",
|
|
62
|
+
"Cosmology",
|
|
63
|
+
"ScipyBackground",
|
|
64
|
+
"JaxBackground",
|
|
65
|
+
"frozen_field",
|
|
66
|
+
"HorndeskiODE",
|
|
67
|
+
"HiClassState",
|
|
68
|
+
"ScipySolver",
|
|
69
|
+
"JAXSolver",
|
|
70
|
+
"get_symbols",
|
|
71
|
+
"SymbolTable",
|
|
72
|
+
"SingularSystemError",
|
|
73
|
+
]
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def __getattr__(name: str):
|
|
77
|
+
"""Load `JAXSolver` on first access (PEP 562), so `import eft4py` never imports JAX."""
|
|
78
|
+
if name == "JAXSolver":
|
|
79
|
+
from .solvers.strategies.jax import JAXSolver
|
|
80
|
+
|
|
81
|
+
return JAXSolver
|
|
82
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def __dir__() -> list[str]:
|
|
86
|
+
"""List the public names, including the lazily loaded `JAXSolver`."""
|
|
87
|
+
return sorted({*__all__, "DOCS_URL", "main"})
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def main() -> None:
|
|
91
|
+
"""Entry point of the `eft4py` command: print the version and where to find the docs."""
|
|
92
|
+
print(f"eft4py {__version__}")
|
|
93
|
+
print("Horndeski dark-energy backgrounds in Python, validated against hi_class.")
|
|
94
|
+
print(f"Documentation: {DOCS_URL}")
|
|
@@ -0,0 +1,138 @@
|
|
|
1
|
+
r"""Numerical backends: how compiled expressions are evaluated.
|
|
2
|
+
|
|
3
|
+
eft4py compiles its symbolic equations with SymPy's `lambdify`, and the same
|
|
4
|
+
expressions can be evaluated in three ways:
|
|
5
|
+
|
|
6
|
+
- `"math"`: plain Python floats. Fastest for one state at a time, which is
|
|
7
|
+
what SciPy's integration loop needs.
|
|
8
|
+
- `"numpy"`: element-wise on NumPy arrays, e.g. a whole trajectory at once.
|
|
9
|
+
- `"jax"`: `jax.numpy`, traceable by `jax.jit`, `vmap` and `grad`.
|
|
10
|
+
|
|
11
|
+
A [`Backend`][eft4py.backend.Backend] bundles everything that differs between
|
|
12
|
+
them. Most users only pass the backend *name* (e.g.
|
|
13
|
+
`HorndeskiODE(model, backend="jax")`); the class matters when writing new
|
|
14
|
+
compiled components.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from dataclasses import dataclass
|
|
18
|
+
from functools import lru_cache
|
|
19
|
+
from typing import Any, Literal, Sequence
|
|
20
|
+
|
|
21
|
+
import numpy as np
|
|
22
|
+
|
|
23
|
+
BackendName = Literal["math", "numpy", "jax"]
|
|
24
|
+
"""Names of the available numerical backends."""
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass(frozen=True)
|
|
28
|
+
class Backend:
|
|
29
|
+
"""A numerical backend for compiled expressions.
|
|
30
|
+
|
|
31
|
+
Attributes:
|
|
32
|
+
name: Backend name.
|
|
33
|
+
lambdify_modules: The `modules` argument passed to `sympy.lambdify`.
|
|
34
|
+
traceable: Whether values may be abstract tracers (JAX under `jit`),
|
|
35
|
+
so that Python control flow on them is not allowed.
|
|
36
|
+
xp: Array namespace (`numpy` or `jax.numpy`); `None` for `"math"`.
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
name: BackendName
|
|
40
|
+
lambdify_modules: Any
|
|
41
|
+
traceable: bool
|
|
42
|
+
xp: Any
|
|
43
|
+
|
|
44
|
+
def where(self, cond: Any, x: Any, y: Any) -> Any:
|
|
45
|
+
"""Element-wise `x if cond else y`, valid for this backend's values."""
|
|
46
|
+
if self.xp is None:
|
|
47
|
+
return x if cond else y
|
|
48
|
+
return self.xp.where(cond, x, y)
|
|
49
|
+
|
|
50
|
+
def broadcast(self, values: Sequence[Any], inputs: Sequence[Any]) -> list:
|
|
51
|
+
r"""Give every output the broadcast shape of the inputs.
|
|
52
|
+
|
|
53
|
+
`lambdify` returns a plain scalar for an expression that simplifies
|
|
54
|
+
to a constant (e.g. $\alpha_T = 0$ for quintessence), even when the
|
|
55
|
+
inputs are arrays. On the `"math"` backend the values are returned
|
|
56
|
+
unchanged.
|
|
57
|
+
|
|
58
|
+
Args:
|
|
59
|
+
values: Outputs of a compiled function.
|
|
60
|
+
inputs: The arguments it was called with.
|
|
61
|
+
|
|
62
|
+
Returns:
|
|
63
|
+
The outputs as arrays of the inputs' broadcast shape.
|
|
64
|
+
"""
|
|
65
|
+
if self.xp is None:
|
|
66
|
+
return list(values)
|
|
67
|
+
shape = np.broadcast_shapes(*(np.shape(x) for x in inputs))
|
|
68
|
+
zeros = self.xp.zeros(shape)
|
|
69
|
+
return [zeros + v for v in values]
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
@lru_cache(maxsize=None)
|
|
73
|
+
def _build(name: str) -> Backend:
|
|
74
|
+
if name == "math":
|
|
75
|
+
return Backend(name="math", lambdify_modules="math", traceable=False, xp=None)
|
|
76
|
+
if name == "numpy":
|
|
77
|
+
return Backend(name="numpy", lambdify_modules="numpy", traceable=False, xp=np)
|
|
78
|
+
if name == "jax":
|
|
79
|
+
import jax.numpy as jnp
|
|
80
|
+
|
|
81
|
+
modules = [
|
|
82
|
+
{
|
|
83
|
+
"exp": jnp.exp, "log": jnp.log, "sqrt": jnp.sqrt,
|
|
84
|
+
"Abs": jnp.abs, "sign": jnp.sign,
|
|
85
|
+
"sin": jnp.sin, "cos": jnp.cos, "tan": jnp.tan,
|
|
86
|
+
},
|
|
87
|
+
jnp,
|
|
88
|
+
]
|
|
89
|
+
return Backend(name="jax", lambdify_modules=modules, traceable=True, xp=jnp)
|
|
90
|
+
raise ValueError(f"Unknown backend {name!r}; expected 'math', 'numpy' or 'jax'.")
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def get_backend(backend: "BackendName | Backend") -> Backend:
|
|
94
|
+
"""Resolve a backend name (or pass a `Backend` through).
|
|
95
|
+
|
|
96
|
+
JAX is imported only when the `"jax"` backend is requested. Backends are
|
|
97
|
+
built once and cached.
|
|
98
|
+
|
|
99
|
+
Args:
|
|
100
|
+
backend: `"math"`, `"numpy"`, `"jax"`, or a `Backend` instance.
|
|
101
|
+
|
|
102
|
+
Returns:
|
|
103
|
+
The corresponding `Backend`.
|
|
104
|
+
|
|
105
|
+
Raises:
|
|
106
|
+
ValueError: If the name is unknown.
|
|
107
|
+
ImportError: If `"jax"` is requested but JAX is not installed.
|
|
108
|
+
"""
|
|
109
|
+
if isinstance(backend, Backend):
|
|
110
|
+
return backend
|
|
111
|
+
return _build(backend)
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
def namespace_of(array: Any) -> Backend:
|
|
115
|
+
"""The backend that an array belongs to: `"numpy"` or `"jax"`.
|
|
116
|
+
|
|
117
|
+
Lets a function accept NumPy or JAX data without a `backend` argument.
|
|
118
|
+
JAX is imported only when the array is not a NumPy array.
|
|
119
|
+
|
|
120
|
+
Args:
|
|
121
|
+
array: A computed array, e.g. a field of a solved background.
|
|
122
|
+
|
|
123
|
+
Returns:
|
|
124
|
+
The `"numpy"` backend for NumPy arrays and scalars and Python
|
|
125
|
+
numbers, the `"jax"` backend otherwise (JAX arrays and tracers inside
|
|
126
|
+
`jit`, `vmap` or `grad`).
|
|
127
|
+
|
|
128
|
+
Warning:
|
|
129
|
+
Pass a *computed* field (such as `H`), not an input passed through
|
|
130
|
+
unchanged (such as `ln_a`): `jax.jit` forwards an input it returns
|
|
131
|
+
unchanged, so a NumPy grid given to a jitted solve comes back as NumPy
|
|
132
|
+
even though everything computed from it is JAX.
|
|
133
|
+
"""
|
|
134
|
+
is_numpy = isinstance(array, (np.ndarray, np.generic, float, int))
|
|
135
|
+
return _build("numpy" if is_numpy else "jax")
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
__all__ = ["Backend", "BackendName", "get_backend", "namespace_of"]
|
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
"""Core types and protocols for eft4py.
|
|
2
|
+
|
|
3
|
+
- Types: the background state `HiClassState` (with the kinetic term $X$
|
|
4
|
+
always derived from it), `HiClassDerivatives`, and the generic
|
|
5
|
+
`ParamPoint`.
|
|
6
|
+
- Protocols: the structural contracts `StateVector`, `ODESystem` and
|
|
7
|
+
`IntegrationStrategy` for extending eft4py.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from .protocols import (
|
|
11
|
+
StateVector,
|
|
12
|
+
ODESystem,
|
|
13
|
+
IntegrationStrategy,
|
|
14
|
+
)
|
|
15
|
+
from .types import (
|
|
16
|
+
ParamPoint,
|
|
17
|
+
HiClassState,
|
|
18
|
+
HiClassDerivatives,
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
__all__ = [
|
|
22
|
+
# Protocols
|
|
23
|
+
"StateVector",
|
|
24
|
+
"ODESystem",
|
|
25
|
+
"IntegrationStrategy",
|
|
26
|
+
# Types
|
|
27
|
+
"ParamPoint",
|
|
28
|
+
"HiClassState",
|
|
29
|
+
"HiClassDerivatives",
|
|
30
|
+
]
|
|
@@ -0,0 +1,167 @@
|
|
|
1
|
+
"""Protocol definitions for abstract interfaces in eft4py.
|
|
2
|
+
|
|
3
|
+
Protocols define the structural contracts for extending eft4py: a state
|
|
4
|
+
(`StateVector`), an ODE system (`ODESystem`) and an integration strategy
|
|
5
|
+
(`IntegrationStrategy`). The base classes `ODESystemBase` and
|
|
6
|
+
`IntegrationStrategyBase` implement them.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from typing import Protocol, Any, Mapping
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class StateVector(Protocol):
|
|
13
|
+
r"""Protocol for any point in configuration space.
|
|
14
|
+
|
|
15
|
+
A StateVector represents a point in the system's phase space. It can be:
|
|
16
|
+
|
|
17
|
+
- A cosmological background state ($\tau$, $a$, $H$, $\phi$, $\phi'$)
|
|
18
|
+
- A parameter evaluation point ($\phi$, $X$, $H$, $M_{\rm pl}$)
|
|
19
|
+
- Any other collection of named variables
|
|
20
|
+
|
|
21
|
+
Key properties:
|
|
22
|
+
|
|
23
|
+
- All values are float (can be wrapped in arrays for batch operations)
|
|
24
|
+
- Keys are canonical symbol names (strings)
|
|
25
|
+
- Used for both symbolic and numerical contexts
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
@property
|
|
29
|
+
def values(self) -> Mapping[str, float]:
|
|
30
|
+
"""Return dict-like mapping of all symbol values.
|
|
31
|
+
|
|
32
|
+
Returns:
|
|
33
|
+
Mapping where keys are symbol names (e.g., 'phi', 'X', 'H')
|
|
34
|
+
and values are float numbers.
|
|
35
|
+
"""
|
|
36
|
+
...
|
|
37
|
+
|
|
38
|
+
def __getitem__(self, key: str) -> float:
|
|
39
|
+
"""Get value by symbol name (convenience method).
|
|
40
|
+
|
|
41
|
+
Args:
|
|
42
|
+
key: Symbol name as string
|
|
43
|
+
|
|
44
|
+
Returns:
|
|
45
|
+
Numerical value
|
|
46
|
+
"""
|
|
47
|
+
...
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class ODESystem(Protocol):
|
|
51
|
+
r"""Protocol for ODE systems of form $dy/d\tau = f(y, \tau)$.
|
|
52
|
+
|
|
53
|
+
An ODESystem encapsulates a complete system of differential equations.
|
|
54
|
+
It handles:
|
|
55
|
+
|
|
56
|
+
- Computing derivatives at any state point
|
|
57
|
+
- Describing the state space dimension
|
|
58
|
+
- Validating initial conditions
|
|
59
|
+
|
|
60
|
+
In eft4py, `HorndeskiODE` implements this protocol for Horndeski
|
|
61
|
+
background evolution in conformal time $\tau$; the protocol itself is
|
|
62
|
+
general.
|
|
63
|
+
|
|
64
|
+
State for Horndeski background: $[a, H, \phi, \phi']$
|
|
65
|
+
Integration variable: $\tau$ (conformal time)
|
|
66
|
+
"""
|
|
67
|
+
|
|
68
|
+
def state_dimension(self) -> int:
|
|
69
|
+
r"""Dimension of state vector.
|
|
70
|
+
|
|
71
|
+
Returns:
|
|
72
|
+
Integer number of state variables
|
|
73
|
+
|
|
74
|
+
Example:
|
|
75
|
+
Horndeski background in conformal time has 4 variables:
|
|
76
|
+
$a, H, \phi, \phi'$
|
|
77
|
+
"""
|
|
78
|
+
...
|
|
79
|
+
|
|
80
|
+
def compute_derivatives(self, state: StateVector, tau: float) -> list[float]:
|
|
81
|
+
r"""Compute $dy/d\tau$ at given state.
|
|
82
|
+
|
|
83
|
+
Args:
|
|
84
|
+
state: Current state with all required variables
|
|
85
|
+
tau: Current conformal time (for reference/diagnostics)
|
|
86
|
+
|
|
87
|
+
Returns:
|
|
88
|
+
List of derivatives $[da/d\tau, dH/d\tau, d\phi/d\tau, d\phi'/d\tau]$
|
|
89
|
+
in correct order
|
|
90
|
+
|
|
91
|
+
Raises:
|
|
92
|
+
ValueError: If state is invalid or computation fails
|
|
93
|
+
"""
|
|
94
|
+
...
|
|
95
|
+
|
|
96
|
+
def validate_state(self, state: StateVector) -> bool:
|
|
97
|
+
r"""Check if state is physically valid.
|
|
98
|
+
|
|
99
|
+
Args:
|
|
100
|
+
state: State to validate
|
|
101
|
+
|
|
102
|
+
Returns:
|
|
103
|
+
True if valid, False otherwise
|
|
104
|
+
|
|
105
|
+
Example:
|
|
106
|
+
For background evolution, might check $a > 0$, $M_*^2 > 0$
|
|
107
|
+
"""
|
|
108
|
+
...
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
class IntegrationStrategy(Protocol):
|
|
112
|
+
"""Protocol for strategies that integrate ODE systems.
|
|
113
|
+
|
|
114
|
+
An IntegrationStrategy is a backend for solving ODE initial value problems.
|
|
115
|
+
Different strategies can use different algorithms:
|
|
116
|
+
|
|
117
|
+
- SciPy's RK45 (robust, variable step)
|
|
118
|
+
- JAX's diffrax (differentiable)
|
|
119
|
+
- Custom RK4 (simple, predictable)
|
|
120
|
+
|
|
121
|
+
All strategies present the same interface for pluggability.
|
|
122
|
+
"""
|
|
123
|
+
|
|
124
|
+
def integrate(
|
|
125
|
+
self,
|
|
126
|
+
system: ODESystem,
|
|
127
|
+
initial_state: StateVector,
|
|
128
|
+
tau_span: tuple[float, float],
|
|
129
|
+
tau_eval: list[float] | None = None,
|
|
130
|
+
**kwargs: Any
|
|
131
|
+
) -> list[StateVector]:
|
|
132
|
+
r"""Integrate ODE system over conformal time interval.
|
|
133
|
+
|
|
134
|
+
Args:
|
|
135
|
+
system: ODESystem to integrate
|
|
136
|
+
initial_state: Initial condition at $\tau_{\rm start}$
|
|
137
|
+
tau_span: $(\tau_{\rm start}, \tau_{\rm end})$ conformal time interval
|
|
138
|
+
tau_eval: Optional list of $\tau$ values to evaluate solution at.
|
|
139
|
+
If None, solver chooses sampling.
|
|
140
|
+
**kwargs: Backend-specific options (rtol, atol, max_steps, etc.)
|
|
141
|
+
|
|
142
|
+
Returns:
|
|
143
|
+
List of StateVector at requested times (or solver-chosen times)
|
|
144
|
+
|
|
145
|
+
Raises:
|
|
146
|
+
RuntimeError: If integration fails (NaN, divergence, etc.)
|
|
147
|
+
ValueError: If invalid parameters
|
|
148
|
+
|
|
149
|
+
Note:
|
|
150
|
+
Results should be in increasing order of $\tau$ (or conformal time).
|
|
151
|
+
"""
|
|
152
|
+
...
|
|
153
|
+
|
|
154
|
+
def name(self) -> str:
|
|
155
|
+
"""Return strategy identifier.
|
|
156
|
+
|
|
157
|
+
Returns:
|
|
158
|
+
String like 'scipy', 'jax', 'rk4'
|
|
159
|
+
"""
|
|
160
|
+
...
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
__all__ = [
|
|
164
|
+
"StateVector",
|
|
165
|
+
"ODESystem",
|
|
166
|
+
"IntegrationStrategy",
|
|
167
|
+
]
|
|
@@ -0,0 +1,175 @@
|
|
|
1
|
+
r"""Type definitions following hi_class conventions.
|
|
2
|
+
|
|
3
|
+
This module defines the core data types used throughout eft4py:
|
|
4
|
+
|
|
5
|
+
- ParamPoint: Evaluation point in parameter space ($\phi$, $X$, $H$, etc.)
|
|
6
|
+
- HiClassState: Point in background evolution following hi_class conventions
|
|
7
|
+
- HiClassDerivatives: Derivatives at a state point
|
|
8
|
+
|
|
9
|
+
$X = \tfrac{1}{2}\dot\phi^2$ is always derived from the state on
|
|
10
|
+
HiClassState, never stored or set independently, so it is guaranteed to
|
|
11
|
+
stay consistent with $\phi'$ and $a$.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from dataclasses import dataclass
|
|
15
|
+
from typing import Mapping
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@dataclass(frozen=True)
|
|
19
|
+
class ParamPoint:
|
|
20
|
+
"""Evaluation point in parameter space.
|
|
21
|
+
|
|
22
|
+
A ParamPoint represents values of symbols at a specific location,
|
|
23
|
+
used for evaluating arbitrary symbolic functions.
|
|
24
|
+
|
|
25
|
+
Examples:
|
|
26
|
+
>>> ParamPoint(values={"phi": 1.5, "X": 0.2, "H": 0.1, "Mpl": 1.0})
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
values: dict[str, float]
|
|
30
|
+
|
|
31
|
+
def __getitem__(self, key: str) -> float:
|
|
32
|
+
"""Get value by symbol name."""
|
|
33
|
+
return self.values[key]
|
|
34
|
+
|
|
35
|
+
def __contains__(self, key: str) -> bool:
|
|
36
|
+
"""Check if symbol is present."""
|
|
37
|
+
return key in self.values
|
|
38
|
+
|
|
39
|
+
@property
|
|
40
|
+
def symbols(self) -> Mapping[str, float]:
|
|
41
|
+
"""Satisfy StateVector protocol."""
|
|
42
|
+
return self.values
|
|
43
|
+
|
|
44
|
+
def get(self, key: str, default: float | None = None) -> float | None:
|
|
45
|
+
"""Get value with optional default."""
|
|
46
|
+
return self.values.get(key, default)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
@dataclass(frozen=True)
|
|
50
|
+
class HiClassState:
|
|
51
|
+
r"""Background evolution state following hi_class conventions.
|
|
52
|
+
|
|
53
|
+
Integration variable: $\tau$ (conformal time)
|
|
54
|
+
State variables: $[a, H, \phi, \phi']$
|
|
55
|
+
|
|
56
|
+
Relationships:
|
|
57
|
+
|
|
58
|
+
- $a$: scale factor
|
|
59
|
+
- $H = a'/a^2$ (conformal Hubble parameter, where $' = d/d\tau$)
|
|
60
|
+
- $\phi$: scalar field
|
|
61
|
+
- $\phi' = d\phi/d\tau$ (conformal time derivative of scalar field)
|
|
62
|
+
|
|
63
|
+
DERIVED (not independent):
|
|
64
|
+
|
|
65
|
+
- $\dot\phi = \phi'/a$ (physical time derivative)
|
|
66
|
+
- $X = \tfrac{1}{2}\dot\phi^2 = \tfrac{1}{2}(\phi'/a)^2$ (kinetic term)
|
|
67
|
+
|
|
68
|
+
$X$ and $\dot\phi$ are always computed from $\phi'$ and $a$, so they can
|
|
69
|
+
never be set independently or become inconsistent with each other.
|
|
70
|
+
"""
|
|
71
|
+
|
|
72
|
+
tau: float # Current conformal time
|
|
73
|
+
a: float # Scale factor
|
|
74
|
+
H: float # Conformal Hubble parameter H = a'/a²
|
|
75
|
+
phi: float # Scalar field value
|
|
76
|
+
phi_prime: float # Conformal time derivative dφ/dτ
|
|
77
|
+
rho_m: float = 0.0 # Matter energy density ρ_m (pressureless dust)
|
|
78
|
+
rho_r: float = 0.0 # Radiation energy density ρ_r (p_r = ρ_r/3)
|
|
79
|
+
|
|
80
|
+
@property
|
|
81
|
+
def phidot(self) -> float:
|
|
82
|
+
r"""Physical time derivative: $\dot\phi = \phi'/a$.
|
|
83
|
+
|
|
84
|
+
Returns:
|
|
85
|
+
Physical scalar field velocity.
|
|
86
|
+
|
|
87
|
+
Note:
|
|
88
|
+
This is derived from state variables (not independent),
|
|
89
|
+
ensuring consistency with X.
|
|
90
|
+
"""
|
|
91
|
+
return self.phi_prime / self.a
|
|
92
|
+
|
|
93
|
+
@property
|
|
94
|
+
def X(self) -> float:
|
|
95
|
+
r"""Kinetic term $X = \tfrac{1}{2}\dot\phi^2$ (ENFORCED - never independent!).
|
|
96
|
+
|
|
97
|
+
Returns:
|
|
98
|
+
$X = \tfrac{1}{2}\dot\phi^2 = \tfrac{1}{2}(\phi'/a)^2$.
|
|
99
|
+
|
|
100
|
+
Warning:
|
|
101
|
+
This property is always computed from state variables. There
|
|
102
|
+
is no way to set X independently, which keeps X and
|
|
103
|
+
$\dot\phi$ consistent by construction.
|
|
104
|
+
"""
|
|
105
|
+
return 0.5 * self.phidot**2
|
|
106
|
+
|
|
107
|
+
@property
|
|
108
|
+
def z(self) -> float:
|
|
109
|
+
"""Redshift z = 1/a - 1 (derived for convenience).
|
|
110
|
+
|
|
111
|
+
Returns:
|
|
112
|
+
Redshift parameter.
|
|
113
|
+
"""
|
|
114
|
+
return 1.0 / self.a - 1.0
|
|
115
|
+
|
|
116
|
+
@property
|
|
117
|
+
def values(self) -> Mapping[str, float]:
|
|
118
|
+
"""Satisfy StateVector protocol.
|
|
119
|
+
|
|
120
|
+
Returns a mapping of all state variables by canonical names.
|
|
121
|
+
Note: X is computed dynamically (not stored).
|
|
122
|
+
"""
|
|
123
|
+
return {
|
|
124
|
+
"tau": self.tau,
|
|
125
|
+
"a": self.a,
|
|
126
|
+
"H": self.H,
|
|
127
|
+
"phi": self.phi,
|
|
128
|
+
"phi_prime": self.phi_prime,
|
|
129
|
+
"phidot": self.phidot,
|
|
130
|
+
"X": self.X,
|
|
131
|
+
"z": self.z,
|
|
132
|
+
"rho_m": self.rho_m,
|
|
133
|
+
"rho_r": self.rho_r,
|
|
134
|
+
}
|
|
135
|
+
|
|
136
|
+
def __getitem__(self, key: str) -> float:
|
|
137
|
+
"""Get value by symbol name (StateVector protocol)."""
|
|
138
|
+
return self.values[key]
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
@dataclass(frozen=True)
|
|
142
|
+
class HiClassDerivatives:
|
|
143
|
+
r"""Derivatives of state variables with respect to conformal time $\tau$.
|
|
144
|
+
|
|
145
|
+
Contains: $[da/d\tau, dH/d\tau, d\phi/d\tau, d\phi'/d\tau, d\rho_m/d\tau, d\rho_r/d\tau]$
|
|
146
|
+
|
|
147
|
+
This is the output of ODESystem.compute_derivatives().
|
|
148
|
+
|
|
149
|
+
Note:
|
|
150
|
+
All derivatives are with respect to conformal time $\tau$, using
|
|
151
|
+
hi_class notation (prime = $d/d\tau$).
|
|
152
|
+
"""
|
|
153
|
+
|
|
154
|
+
da_dtau: float # da/dτ (scale factor derivative)
|
|
155
|
+
dH_dtau: float # dH/dτ (Hubble derivative)
|
|
156
|
+
dphi_dtau: float # dφ/dτ (scalar field derivative)
|
|
157
|
+
dphi_prime_dtau: float # dφ'/dτ (acceleration in conformal time)
|
|
158
|
+
drho_m_dtau: float = 0.0 # dρ_m/dτ (dust continuity: -3aH·ρ_m)
|
|
159
|
+
drho_r_dtau: float = 0.0 # dρ_r/dτ (radiation continuity: -4aH·ρ_r)
|
|
160
|
+
|
|
161
|
+
def to_list(self) -> list[float]:
|
|
162
|
+
r"""Convert to ordered list for ODE solver.
|
|
163
|
+
|
|
164
|
+
Returns:
|
|
165
|
+
$[da/d\tau, dH/d\tau, d\phi/d\tau, d\phi'/d\tau, d\rho_m/d\tau, d\rho_r/d\tau]$.
|
|
166
|
+
"""
|
|
167
|
+
return [self.da_dtau, self.dH_dtau, self.dphi_dtau, self.dphi_prime_dtau,
|
|
168
|
+
self.drho_m_dtau, self.drho_r_dtau]
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
__all__ = [
|
|
172
|
+
"ParamPoint",
|
|
173
|
+
"HiClassState",
|
|
174
|
+
"HiClassDerivatives",
|
|
175
|
+
]
|