Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,6 @@ classifiers = [

dependencies = [
"absl-py>=2.0.0",
"typing_extensions>=4.2.0",
"immutabledict>=1.0.0",
"jax>=0.11.1",
"jaxlib>=0.11.1",
Expand Down
7 changes: 4 additions & 3 deletions torax/_src/array_typing.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,8 @@
# ============================================================================
"""Common types for using jaxtyping in TORAX."""

from typing import TypeAlias
from collections.abc import Callable
from typing import Any, TypeAlias
import jax
import jaxtyping as jt
import numpy as np
Expand All @@ -38,7 +39,7 @@
BoolVectorFace: TypeAlias = jt.Bool[Array, "rhon+1"]


def jaxtyped[T](fn: T) -> T:
def jaxtyped[T: Callable[..., Any] | type](fn: T) -> T:
"""Function and dataclass decorator to perform runtime type-checking.

This will perform jaxtyping runtime type checking if the environment variable
Expand All @@ -52,6 +53,6 @@ def jaxtyped[T](fn: T) -> T:
"""
runtime_checking = jax_utils.env_bool(name="TORAX_JAXTYPING", default=False)
if runtime_checking:
return jt.jaxtyped(fn, typechecker=typeguard.typechecked) # pyrefly: ignore[no-matching-overload]
return jt.jaxtyped(typechecker=typeguard.typechecked)(fn)
else:
return fn
13 changes: 6 additions & 7 deletions torax/_src/config/build_runtime_params.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@
"""

import dataclasses
from typing import Any, Callable, Mapping, Sequence, TypeAlias
from typing import Any, Callable, Mapping, Self, Sequence, TypeAlias

import chex
import equinox as eqx
Expand All @@ -49,7 +49,6 @@
from torax._src.torax_pydantic import interpolated_param_2d
from torax._src.torax_pydantic import model_config
from torax._src.transport_model import pydantic_model as transport_pydantic_model
import typing_extensions

# pylint: disable=invalid-name

Expand Down Expand Up @@ -97,7 +96,7 @@ class interpolates any time-dependent params in the input config to the values
def from_config(
cls,
config: model_config.ToraxConfig,
) -> typing_extensions.Self:
) -> Self:
"""Constructs a RuntimeParamsProvider from a ToraxConfig."""
return cls(
sources=config.sources,
Expand Down Expand Up @@ -141,11 +140,11 @@ def __call__(
def update_provider(
self,
get_nodes_to_replace: Callable[
[typing_extensions.Self],
[Self],
Sequence[ReplaceablePytreeNodes],
],
replacement_values: Sequence[ValidUpdates],
) -> typing_extensions.Self:
) -> Self:
"""Updates a provider with new values. Works under `jax.jit`.

Example usage:
Expand Down Expand Up @@ -205,7 +204,7 @@ def get_node_from_path(self, path: str) -> Any:

def update_provider_from_mapping(
self, replacements: Mapping[str, ValidUpdates]
) -> typing_extensions.Self:
) -> Self:
"""Update a provider from a mapping of replacements.

Example usage:
Expand Down Expand Up @@ -238,7 +237,7 @@ def update_provider_from_mapping(
"""

def get_replacements(
provider: typing_extensions.Self,
provider: Self,
) -> list[ReplaceablePytreeNodes]:
"""Returns the nodes to replace."""
nodes_to_replace: list[ReplaceablePytreeNodes] = []
Expand Down
3 changes: 1 addition & 2 deletions torax/_src/config/numerics.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,14 +16,13 @@

import dataclasses
import functools
from typing import Annotated
from typing import Annotated, Self

import chex
import jax
import pydantic
from torax._src import array_typing
from torax._src.torax_pydantic import torax_pydantic
from typing_extensions import Self


# pylint: disable=invalid-name
Expand Down
34 changes: 17 additions & 17 deletions torax/_src/core_profiles/getters.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,7 +80,7 @@ def get_updated_ion_temperature(
face_centers=geo.rho_face_norm,
left_face_grad_constraint=jnp.zeros(()),
right_face_grad_constraint=None,
right_face_constraint=profile_conditions_params.T_i_right_bc, # pyrefly: ignore[bad-argument-type]
right_face_constraint=profile_conditions_params.T_i_right_bc,
)
return T_i

Expand All @@ -107,7 +107,7 @@ def get_updated_electron_temperature(
face_centers=geo.rho_face_norm,
left_face_grad_constraint=jnp.zeros(()),
right_face_grad_constraint=None,
right_face_constraint=profile_conditions_params.T_e_right_bc, # pyrefly: ignore[bad-argument-type]
right_face_constraint=profile_conditions_params.T_e_right_bc,
)
return T_e

Expand Down Expand Up @@ -261,7 +261,7 @@ def get_updated_toroidal_angular_velocity(
value=value,
face_centers=geo.rho_face_norm,
right_face_grad_constraint=None,
right_face_constraint=profile_conditions_params.toroidal_angular_velocity_right_bc, # pyrefly: ignore[bad-argument-type]
right_face_constraint=profile_conditions_params.toroidal_angular_velocity_right_bc,
)
return toroidal_angular_velocity

Expand Down Expand Up @@ -293,14 +293,14 @@ def _get_ion_properties_from_fractions(
"""Calculates ion properties when impurity content is defined by fractions."""

charge_state_info = charge_states.get_average_charge_state(
T_e=T_e.value, # pyrefly: ignore[bad-argument-type]
T_e=T_e.value,
fractions=impurity_params.fractions,
Z_override=impurity_params.Z_override,
)
Z_impurity = charge_state_info.Z_mixture

charge_state_info_face = charge_states.get_average_charge_state(
T_e=T_e.face_value(), # pyrefly: ignore[bad-argument-type]
T_e=T_e.face_value(),
fractions=impurity_params.fractions_face,
Z_override=impurity_params.Z_override,
)
Expand Down Expand Up @@ -343,12 +343,12 @@ def _get_ion_properties_from_n_e_ratios(
) -> _IonProperties:
"""Calculates ion properties when impurity content is defined by n_e ratios."""
average_charge_state = charge_states.get_average_charge_state(
T_e=T_e.value, # pyrefly: ignore[bad-argument-type]
T_e=T_e.value,
fractions=impurity_params.fractions,
Z_override=impurity_params.Z_override,
)
average_charge_state_face = charge_states.get_average_charge_state(
T_e=T_e.face_value(), # pyrefly: ignore[bad-argument-type]
T_e=T_e.face_value(),
fractions=impurity_params.fractions_face,
Z_override=impurity_params.Z_override,
)
Expand Down Expand Up @@ -439,13 +439,13 @@ def _get_ion_properties_from_n_e_ratios_Z_eff(
impurity_symbols = tuple(impurity_params.n_e_ratios.keys())
Z_per_species = jnp.stack([
charge_states.calculate_average_charge_state_single_species(
T_e.value, symbol # pyrefly: ignore[bad-argument-type]
T_e.value, symbol
)
for symbol in impurity_symbols
])
Z_per_species_face = jnp.stack([
charge_states.calculate_average_charge_state_single_species(
T_e.face_value(), symbol # pyrefly: ignore[bad-argument-type]
T_e.face_value(), symbol
)
for symbol in impurity_symbols
])
Expand Down Expand Up @@ -561,14 +561,14 @@ def _solve_system(a1, a2, b1, b2, c1, c2):
)

charge_state_info = charge_states.get_average_charge_state(
T_e=T_e.value, # pyrefly: ignore[bad-argument-type]
T_e=T_e.value,
fractions=fractions, # pyrefly: ignore[bad-argument-type]
Z_override=impurity_params.Z_override,
)
Z_impurity = charge_state_info.Z_mixture

charge_state_info_face = charge_states.get_average_charge_state(
T_e=T_e.face_value(), # pyrefly: ignore[bad-argument-type]
T_e=T_e.face_value(),
fractions=fractions_face, # pyrefly: ignore[bad-argument-type]
Z_override=impurity_params.Z_override,
)
Expand Down Expand Up @@ -631,12 +631,12 @@ def get_updated_ions(
"""

Z_i = charge_states.get_average_charge_state(
T_e=T_e.value, # pyrefly: ignore[bad-argument-type]
T_e=T_e.value,
fractions=runtime_params.plasma_composition.main_ion.fractions, # pyrefly: ignore[bad-argument-type]
Z_override=runtime_params.plasma_composition.main_ion.Z_override,
).Z_mixture
Z_i_face = charge_states.get_average_charge_state(
T_e=T_e.face_value(), # pyrefly: ignore[bad-argument-type]
T_e=T_e.face_value(),
fractions=runtime_params.plasma_composition.main_ion.fractions, # pyrefly: ignore[bad-argument-type]
Z_override=runtime_params.plasma_composition.main_ion.Z_override,
).Z_mixture
Expand Down Expand Up @@ -678,7 +678,7 @@ def get_updated_ions(
value=n_e.value * ion_properties.dilution_factor,
face_centers=geo.rho_face_norm,
right_face_grad_constraint=None,
right_face_constraint=n_e.right_face_constraint # pyrefly: ignore[bad-argument-type, unsupported-operation]
right_face_constraint=n_e.right_face_constraint # pyrefly: ignore[unsupported-operation]
* ion_properties.dilution_factor_edge,
)

Expand Down Expand Up @@ -724,9 +724,9 @@ def get_updated_ions(
Z_eff_face = _calculate_Z_eff(
Z_i_face,
ion_properties.Z_impurity_face,
n_i.face_value(), # pyrefly: ignore[bad-argument-type]
n_impurity.face_value(), # pyrefly: ignore[bad-argument-type]
n_e.face_value(), # pyrefly: ignore[bad-argument-type]
n_i.face_value(),
n_impurity.face_value(),
n_e.face_value(),
)

# Convert array of fractions to a mapping from symbol to fraction profile.
Expand Down
2 changes: 1 addition & 1 deletion torax/_src/core_profiles/initialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -529,7 +529,7 @@ def _calculate_all_psi_dependent_profiles(
psidot = dataclasses.replace(
core_profiles.psidot,
value=psidot_value,
right_face_constraint=v_loop_lcfs, # pyrefly: ignore[bad-argument-type]
right_face_constraint=v_loop_lcfs,
right_face_grad_constraint=None,
)
core_profiles = dataclasses.replace(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
"""Impurity content defined by ratios of impurity to electron density."""

import dataclasses
from typing import Annotated, Literal, Mapping
from typing import Annotated, Literal, Mapping, Self

import chex
import jax
Expand All @@ -25,15 +25,14 @@
from torax._src import array_typing
from torax._src import constants
from torax._src.torax_pydantic import torax_pydantic
import typing_extensions


# pylint: disable=invalid-name


def calculate_fractions_from_ratios(
ratios: Mapping[str, chex.Array],
) -> Mapping[str, chex.Array]:
def calculate_fractions_from_ratios[ArrayT: array_typing.Array](
ratios: Mapping[str, ArrayT],
) -> Mapping[str, ArrayT]:
"""Calculates fractions from ratios, handling the all-zero case."""
# Ratios can be 1D (n_species,) or 2D (n_species, n_grid).
# Sum over the species axis.
Expand Down Expand Up @@ -72,12 +71,12 @@ class RuntimeParams:
@property
def fractions(self) -> Mapping[str, array_typing.FloatVector]:
"""Returns the impurity fractions calculated from the n_e_ratios."""
return calculate_fractions_from_ratios(self.n_e_ratios) # pyrefly: ignore[bad-return]
return calculate_fractions_from_ratios(self.n_e_ratios)

@property
def fractions_face(self) -> Mapping[str, array_typing.FloatVectorFace]:
"""Returns the impurity fractions calculated from the n_e_ratios."""
return calculate_fractions_from_ratios(self.n_e_ratios_face) # pyrefly: ignore[bad-return]
return calculate_fractions_from_ratios(self.n_e_ratios_face)


class ElectronDensityRatios(torax_pydantic.BaseModelFrozen):
Expand All @@ -99,7 +98,7 @@ class ElectronDensityRatios(torax_pydantic.BaseModelFrozen):
)

@pydantic.model_validator(mode='after')
def _validate_species_not_empty(self) -> typing_extensions.Self:
def _validate_species_not_empty(self) -> Self:
if not self.species:
raise ValueError('The species dictionary cannot be empty.')
return self
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,13 +14,12 @@

"""Impurity content defined by ratios, with one species constrained by Z_eff."""
import dataclasses
from typing import Annotated, Literal, Mapping
from typing import Annotated, Literal, Mapping, Self
import chex
import jax
import pydantic
from torax._src import array_typing
from torax._src.torax_pydantic import torax_pydantic
import typing_extensions

# pylint: disable=invalid-name

Expand Down Expand Up @@ -75,7 +74,7 @@ def build_runtime_params(self, t: chex.Numeric) -> RuntimeParams:
)

@pydantic.model_validator(mode='after')
def _validate_one_none(self) -> typing_extensions.Self:
def _validate_one_none(self) -> Self:
if not self.species:
raise ValueError('The species dictionary cannot be empty.')
none_count = sum(v is None for v in self.species.values())
Expand Down
3 changes: 2 additions & 1 deletion torax/_src/core_profiles/plasma_composition/ion_mixture.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,16 +13,17 @@
# limitations under the License.

"""Ion mixture model and impurity fractions model for plasma composition."""

from collections.abc import Mapping
import dataclasses
from typing import Final
import chex
import jax
from jax import numpy as jnp
from torax._src import array_typing
from torax._src import constants
from torax._src.config import runtime_validation_utils
from torax._src.torax_pydantic import torax_pydantic
from typing_extensions import Final

# pylint: disable=invalid-name
_IMPURITY_MODE_FRACTIONS: Final[str] = 'fractions'
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
import dataclasses
import functools
import logging
from typing import Annotated
from typing import Annotated, Final, Self
import chex
import jax
import numpy as np
Expand All @@ -28,8 +28,6 @@
from torax._src.core_profiles.plasma_composition import impurity_fractions
from torax._src.core_profiles.plasma_composition import ion_mixture
from torax._src.torax_pydantic import torax_pydantic
import typing_extensions
from typing_extensions import Final

# pylint: disable=invalid-name

Expand Down Expand Up @@ -110,7 +108,7 @@ class PlasmaComposition(torax_pydantic.BaseModelFrozen):
A_i_override: torax_pydantic.TimeVaryingScalar | None = None

@pydantic.model_validator(mode='after')
def _check_zeff_usage(self) -> typing_extensions.Self:
def _check_zeff_usage(self) -> Self:
"""Warns user if Z_eff is provided but will be ignored."""
is_default_zeff = all(
np.allclose(val, 1.0) for _, (_, val) in self.Z_eff.value.items()
Expand Down
Loading
Loading