From dd4bad9808a63b6bc4dc5a4a5902684dc0618941 Mon Sep 17 00:00:00 2001 From: Sebastian Bodenstein Date: Mon, 14 Sep 2026 06:02:38 -0700 Subject: [PATCH] Remove pyrefly disables PiperOrigin-RevId: 981104961 --- pyproject.toml | 1 - torax/_src/array_typing.py | 7 +- torax/_src/config/build_runtime_params.py | 13 +-- torax/_src/config/numerics.py | 3 +- torax/_src/core_profiles/getters.py | 34 +++--- torax/_src/core_profiles/initialization.py | 2 +- .../electron_density_ratios.py | 15 ++- .../electron_density_ratios_zeff.py | 5 +- .../plasma_composition/ion_mixture.py | 3 +- .../plasma_composition/plasma_composition.py | 6 +- .../_src/core_profiles/profile_conditions.py | 7 +- .../_src/core_profiles/tests/getters_test.py | 48 ++++---- .../tests/plasma_composition_test.py | 10 +- .../_src/core_profiles/tests/updaters_test.py | 4 +- torax/_src/core_profiles/updaters.py | 2 +- .../edge/extended_lengyel/divertor_sol_1d.py | 34 +++--- .../extended_lengyel_formulas.py | 12 +- .../extended_lengyel_model.py | 8 +- .../extended_lengyel_solvers.py | 4 +- .../edge/extended_lengyel/pydantic_model.py | 9 +- .../tests/extended_lengyel_model_test.py | 4 +- .../tests/extended_lengyel_multistart_test.py | 2 +- .../tests/extended_lengyel_output_test.py | 2 +- torax/_src/fvm/calc_coeffs.py | 13 ++- torax/_src/fvm/cell_variable.py | 4 +- torax/_src/fvm/convection_terms.py | 14 ++- torax/_src/fvm/diffusion_terms.py | 32 ++--- torax/_src/fvm/tests/calc_coeffs_test.py | 2 +- torax/_src/geometry/base.py | 7 +- torax/_src/geometry/chease.py | 5 +- torax/_src/geometry/circular_geometry.py | 6 +- torax/_src/geometry/eqdsk.py | 5 +- torax/_src/geometry/fbt.py | 7 +- torax/_src/geometry/geometry.py | 20 ++-- torax/_src/geometry/geometry_provider.py | 5 +- torax/_src/geometry/imas.py | 5 +- torax/_src/imas_tools/output/core_profiles.py | 8 +- torax/_src/imas_tools/output/equilibrium.py | 4 +- torax/_src/interpolated_param.py | 4 +- torax/_src/math_utils.py | 14 +-- torax/_src/mhd/sawtooth/flatten_profile.py | 4 +- .../sawtooth/tests/flatten_profile_test.py | 18 +-- .../neoclassical/bootstrap_current/redl.py | 8 +- .../neoclassical/bootstrap_current/sauter.py | 8 +- .../_src/neoclassical/conductivity/sauter.py | 4 +- torax/_src/neoclassical/formulas/formulas.py | 4 +- .../formulas/tests/formulas_test.py | 4 +- .../neoclassical/formulas/tests/redl_test.py | 4 +- .../formulas/tests/sauter_test.py | 4 +- .../neoclassical/transport/angioni_sauter.py | 31 +++-- torax/_src/neoclassical/transport/base.py | 4 +- torax/_src/neoclassical/transport/zeros.py | 3 +- torax/_src/orchestration/step_function.py | 4 +- .../orchestration/step_function_processing.py | 25 ++-- torax/_src/output_tools/output.py | 4 +- torax/_src/output_tools/post_processing.py | 27 +++-- .../tests/post_processing_test.py | 26 ++--- .../power_scaling_formation_model.py | 12 +- .../pedestal_model/pedestal_model_output.py | 4 +- torax/_src/pedestal_model/pydantic_model.py | 5 +- .../profile_value_saturation_model.py | 8 +- .../profile_value_saturation_model_test.py | 2 +- .../pedestal_model/set_pped_tpedratio_nped.py | 2 +- torax/_src/pedestal_model/set_tped_nped.py | 2 +- .../tests/pedestal_model_output_test.py | 2 +- torax/_src/physics/collisions.py | 67 +++++------ torax/_src/physics/fast_ion_utils.py | 62 +++++----- torax/_src/physics/formulas.py | 16 +-- torax/_src/physics/psi_calculations.py | 52 +++++---- torax/_src/physics/rotation.py | 4 +- torax/_src/physics/scaling_laws.py | 31 ++--- torax/_src/physics/tests/formulas_test.py | 6 +- .../physics/tests/psi_calculations_test.py | 4 +- .../_src/sources/bremsstrahlung_heat_sink.py | 17 +-- .../sources/cyclotron_radiation_heat_sink.py | 19 ++- .../_src/sources/electron_cyclotron_source.py | 10 +- torax/_src/sources/formulas.py | 22 ++-- torax/_src/sources/fusion_heat_source.py | 6 +- torax/_src/sources/gas_puff_source.py | 8 +- torax/_src/sources/generic_current_source.py | 4 +- .../sources/generic_ion_el_heat_source.py | 4 +- torax/_src/sources/generic_particle_source.py | 10 +- .../impurity_radiation_constant_fraction.py | 2 +- .../impurity_radiation_mavrin_fit.py | 2 +- .../ion_cyclotron_source/scaled_profile.py | 2 +- .../sources/ion_cyclotron_source/toric_nn.py | 87 +++++++------- torax/_src/sources/ohmic_heat_source.py | 6 +- torax/_src/sources/pellet_source.py | 10 +- torax/_src/sources/pydantic_model.py | 3 +- torax/_src/sources/source.py | 1 + torax/_src/sources/source_profiles.py | 15 ++- ...ction_impurity_radiation_heat_sink_test.py | 8 +- .../sources/tests/ohmic_heat_source_test.py | 4 +- .../_src/sources/tests/register_model_test.py | 4 +- .../tests/source_profile_builders_test.py | 6 +- torax/_src/state.py | 109 ++++++++++-------- .../torax_pydantic/interpolated_param_1d.py | 23 ++-- .../torax_pydantic/interpolated_param_2d.py | 25 ++-- torax/_src/torax_pydantic/model_base.py | 15 ++- torax/_src/torax_pydantic/model_config.py | 31 +++-- .../tests/interpolated_param_1d_test.py | 22 ++-- .../tests/interpolated_param_2d_test.py | 12 +- torax/_src/torax_pydantic/torax_pydantic.py | 3 +- torax/_src/transport_model/component.py | 40 ++++--- torax/_src/transport_model/pydantic_model.py | 15 ++- .../transport_model/pydantic_model_base.py | 5 +- torax/_src/transport_model/qlknn_10d.py | 5 +- .../qualikiz_based_transport_model.py | 6 +- .../quasilinear_transport_model.py | 26 +++-- .../qualikiz_based_transport_model_test.py | 7 +- .../tests/quasilinear_transport_model_test.py | 10 +- .../tests/tglf_based_transport_model_test.py | 7 +- .../tglf_based_transport_model.py | 34 +++--- .../transport_coefficients_builder.py | 2 +- .../_src/transport_model/transport_coeffs.py | 6 +- torax/_src/tridiagonal.py | 14 +-- torax/_src/version.py | 3 +- .../experimental/gas_puff_feedback_source.py | 10 +- .../tests/gas_puff_feedback_source_test.py | 2 +- torax/run_simulation_main.py | 20 ++-- torax/tests/sim_time_dependence_test.py | 27 ++++- 121 files changed, 815 insertions(+), 765 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 745aeb980..0524218bf 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/torax/_src/array_typing.py b/torax/_src/array_typing.py index 053ce17eb..4cd3103dc 100644 --- a/torax/_src/array_typing.py +++ b/torax/_src/array_typing.py @@ -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 @@ -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 @@ -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 diff --git a/torax/_src/config/build_runtime_params.py b/torax/_src/config/build_runtime_params.py index 5fc382cc3..aba9678e7 100644 --- a/torax/_src/config/build_runtime_params.py +++ b/torax/_src/config/build_runtime_params.py @@ -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 @@ -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 @@ -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, @@ -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: @@ -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: @@ -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] = [] diff --git a/torax/_src/config/numerics.py b/torax/_src/config/numerics.py index cea199ae7..11bfa9240 100644 --- a/torax/_src/config/numerics.py +++ b/torax/_src/config/numerics.py @@ -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 diff --git a/torax/_src/core_profiles/getters.py b/torax/_src/core_profiles/getters.py index 0ac8f8d28..d8072b6fd 100644 --- a/torax/_src/core_profiles/getters.py +++ b/torax/_src/core_profiles/getters.py @@ -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 @@ -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 @@ -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 @@ -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, ) @@ -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, ) @@ -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 ]) @@ -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, ) @@ -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 @@ -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, ) @@ -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. diff --git a/torax/_src/core_profiles/initialization.py b/torax/_src/core_profiles/initialization.py index 063a2f2fe..1afaf9542 100644 --- a/torax/_src/core_profiles/initialization.py +++ b/torax/_src/core_profiles/initialization.py @@ -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( diff --git a/torax/_src/core_profiles/plasma_composition/electron_density_ratios.py b/torax/_src/core_profiles/plasma_composition/electron_density_ratios.py index 7d491b30c..24e885e4e 100644 --- a/torax/_src/core_profiles/plasma_composition/electron_density_ratios.py +++ b/torax/_src/core_profiles/plasma_composition/electron_density_ratios.py @@ -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 @@ -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. @@ -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): @@ -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 diff --git a/torax/_src/core_profiles/plasma_composition/electron_density_ratios_zeff.py b/torax/_src/core_profiles/plasma_composition/electron_density_ratios_zeff.py index 440c78a38..b973d6602 100644 --- a/torax/_src/core_profiles/plasma_composition/electron_density_ratios_zeff.py +++ b/torax/_src/core_profiles/plasma_composition/electron_density_ratios_zeff.py @@ -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 @@ -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()) diff --git a/torax/_src/core_profiles/plasma_composition/ion_mixture.py b/torax/_src/core_profiles/plasma_composition/ion_mixture.py index 6dc37f7d0..baa206cb7 100644 --- a/torax/_src/core_profiles/plasma_composition/ion_mixture.py +++ b/torax/_src/core_profiles/plasma_composition/ion_mixture.py @@ -13,8 +13,10 @@ # 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 @@ -22,7 +24,6 @@ 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' diff --git a/torax/_src/core_profiles/plasma_composition/plasma_composition.py b/torax/_src/core_profiles/plasma_composition/plasma_composition.py index f367e53e7..060c958c5 100644 --- a/torax/_src/core_profiles/plasma_composition/plasma_composition.py +++ b/torax/_src/core_profiles/plasma_composition/plasma_composition.py @@ -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 @@ -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 @@ -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() diff --git a/torax/_src/core_profiles/profile_conditions.py b/torax/_src/core_profiles/profile_conditions.py index e03af2c93..82902aa02 100644 --- a/torax/_src/core_profiles/profile_conditions.py +++ b/torax/_src/core_profiles/profile_conditions.py @@ -17,7 +17,7 @@ import dataclasses import enum import logging -from typing import Annotated, Callable, Final, Sequence +from typing import Annotated, Callable, Final, Self, Sequence import chex import jax @@ -28,7 +28,6 @@ from torax._src.internal_boundary_conditions import internal_boundary_conditions as internal_boundary_conditions_lib from torax._src.physics import fast_ion as fast_ion_lib from torax._src.torax_pydantic import torax_pydantic -from typing_extensions import Self # pylint: disable=invalid-name @@ -677,13 +676,13 @@ def apply_prescribed_fast_ions( value=p.n, face_centers=face_centers, right_face_grad_constraint=None, - right_face_constraint=p.n_right_bc, # pyrefly: ignore[bad-argument-type] + right_face_constraint=p.n_right_bc, ), T=cell_variable.CellVariable( value=p.T, face_centers=face_centers, right_face_grad_constraint=None, - right_face_constraint=p.T_right_bc, # pyrefly: ignore[bad-argument-type] + right_face_constraint=p.T_right_bc, ), ) ) diff --git a/torax/_src/core_profiles/tests/getters_test.py b/torax/_src/core_profiles/tests/getters_test.py index 4ff6a017b..4f20a8eaa 100644 --- a/torax/_src/core_profiles/tests/getters_test.py +++ b/torax/_src/core_profiles/tests/getters_test.py @@ -333,8 +333,8 @@ def test_n_e_core_profile_setter_with_normalization( ) ratio = n_e_unnormalized.value / n_e_normalized.value - np.all(np.isclose(ratio, ratio[0])) # pyrefly: ignore[bad-index] - self.assertNotEqual(ratio[0], 1.0) # pyrefly: ignore[bad-index] + np.all(np.isclose(ratio, ratio[0])) + self.assertNotEqual(ratio[0], 1.0) @parameterized.parameters( True, @@ -379,8 +379,8 @@ def test_n_e_core_profile_setter_with_fGW( ) ratio = n_e.value / n_e_fGW.value - np.all(np.isclose(ratio, ratio[0])) # pyrefly: ignore[bad-index] - self.assertNotEqual(ratio[0], 1.0) # pyrefly: ignore[bad-index] + np.all(np.isclose(ratio, ratio[0])) + self.assertNotEqual(ratio[0], 1.0) def test_get_updated_ion_data(self): expected_value = np.array([1.4375e20, 1.3125e20, 1.1875e20, 1.0625e20]) @@ -471,17 +471,17 @@ def test_Z_eff_calculation(self): calculated_Z_eff = getters._calculate_Z_eff( core_profiles.Z_i, core_profiles.Z_impurity, - core_profiles.n_i.value, # pyrefly: ignore[bad-argument-type] - core_profiles.n_impurity.value, # pyrefly: ignore[bad-argument-type] - core_profiles.n_e.value, # pyrefly: ignore[bad-argument-type] + core_profiles.n_i.value, + core_profiles.n_impurity.value, + core_profiles.n_e.value, ) calculated_Z_eff_face = getters._calculate_Z_eff( core_profiles.Z_i_face, core_profiles.Z_impurity_face, - core_profiles.n_i.face_value(), # pyrefly: ignore[bad-argument-type] - core_profiles.n_impurity.face_value(), # pyrefly: ignore[bad-argument-type] - core_profiles.n_e.face_value(), # pyrefly: ignore[bad-argument-type] + core_profiles.n_i.face_value(), + core_profiles.n_impurity.face_value(), + core_profiles.n_e.face_value(), ) np.testing.assert_allclose( @@ -569,13 +569,13 @@ def test_get_updated_ions_impurity_mixture(self): T_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, T_e), face_centers=geo.rho_face_norm, - right_face_constraint=T_e, # pyrefly: ignore[bad-argument-type] + right_face_constraint=T_e, right_face_grad_constraint=None, ) n_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, n_e), face_centers=geo.rho_face_norm, - right_face_constraint=n_e, # pyrefly: ignore[bad-argument-type] + right_face_constraint=n_e, right_face_grad_constraint=None, ) ions = getters.get_updated_ions( @@ -674,13 +674,13 @@ def test_get_updated_ions_impurity_mixture_radially_dependent(self): T_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, T_e), face_centers=geo.rho_face_norm, - right_face_constraint=T_e, # pyrefly: ignore[bad-argument-type] + right_face_constraint=T_e, right_face_grad_constraint=None, ) n_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, n_e), face_centers=geo.rho_face_norm, - right_face_constraint=n_e, # pyrefly: ignore[bad-argument-type] + right_face_constraint=n_e, right_face_grad_constraint=None, ) ions = getters.get_updated_ions( @@ -798,13 +798,13 @@ def _run_get_updated_ions(torax_config): t_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, t_e_keV), face_centers=geo.rho_face_norm, - right_face_constraint=t_e_keV, # pyrefly: ignore[bad-argument-type] + right_face_constraint=t_e_keV, right_face_grad_constraint=None, ) n_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, n_e_val), face_centers=geo.rho_face_norm, - right_face_constraint=n_e_val, # pyrefly: ignore[bad-argument-type] + right_face_constraint=n_e_val, right_face_grad_constraint=None, ) return getters.get_updated_ions( @@ -859,13 +859,13 @@ def test_get_updated_ions_with_n_e_ratios_He_main_ion(self): T_e_cv = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, t_e_keV), face_centers=geo.rho_face_norm, - right_face_constraint=t_e_keV, # pyrefly: ignore[bad-argument-type] + right_face_constraint=t_e_keV, right_face_grad_constraint=None, ) n_e_cv = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, n_e_val), face_centers=geo.rho_face_norm, - right_face_constraint=n_e_val, # pyrefly: ignore[bad-argument-type] + right_face_constraint=n_e_val, right_face_grad_constraint=None, ) ions = getters.get_updated_ions( @@ -998,13 +998,13 @@ def _run_get_updated_ions(torax_config): t_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, t_e_keV), face_centers=geo.rho_face_norm, - right_face_constraint=t_e_keV, # pyrefly: ignore[bad-argument-type] + right_face_constraint=t_e_keV, right_face_grad_constraint=None, ) n_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, n_e_val), face_centers=geo.rho_face_norm, - right_face_constraint=n_e_val, # pyrefly: ignore[bad-argument-type] + right_face_constraint=n_e_val, right_face_grad_constraint=None, ) return getters.get_updated_ions( @@ -1068,13 +1068,13 @@ def _run_get_updated_ions(torax_config): t_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, t_e_keV), face_centers=geo.rho_face_norm, - right_face_constraint=t_e_keV, # pyrefly: ignore[bad-argument-type] + right_face_constraint=t_e_keV, right_face_grad_constraint=None, ) n_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, n_e_val), face_centers=geo.rho_face_norm, - right_face_constraint=n_e_val, # pyrefly: ignore[bad-argument-type] + right_face_constraint=n_e_val, right_face_grad_constraint=None, ) return getters.get_updated_ions( @@ -1169,13 +1169,13 @@ def _run_get_updated_ions(torax_config): t_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, t_e_keV), face_centers=geo.rho_face_norm, - right_face_constraint=t_e_keV, # pyrefly: ignore[bad-argument-type] + right_face_constraint=t_e_keV, right_face_grad_constraint=None, ) n_e_cell_variable = cell_variable.CellVariable( value=jnp.full_like(geo.rho_norm, n_e_val), face_centers=geo.rho_face_norm, - right_face_constraint=n_e_val, # pyrefly: ignore[bad-argument-type] + right_face_constraint=n_e_val, right_face_grad_constraint=None, ) return getters.get_updated_ions( diff --git a/torax/_src/core_profiles/tests/plasma_composition_test.py b/torax/_src/core_profiles/tests/plasma_composition_test.py index 1e87bb70e..f732c4b6c 100644 --- a/torax/_src/core_profiles/tests/plasma_composition_test.py +++ b/torax/_src/core_profiles/tests/plasma_composition_test.py @@ -37,10 +37,10 @@ def test_plasma_composition_validation_error_for_unphysical_zeff( self, Z_eff: float ): with self.assertRaises(pydantic.ValidationError): - plasma_composition.PlasmaComposition(Z_eff=Z_eff) # pyrefly: ignore[missing-argument] + plasma_composition.PlasmaComposition(Z_eff=Z_eff) def test_plasma_composition_build_runtime_params_smoke_test(self): - pc = plasma_composition.PlasmaComposition() # pyrefly: ignore[missing-argument] + pc = plasma_composition.PlasmaComposition() geo = circular_geometry.CircularConfig().build_geometry() torax_pydantic.set_grid(pc, geo.torax_mesh) pc.build_runtime_params(t=0.0) @@ -52,7 +52,7 @@ def test_plasma_composition_build_runtime_params_smoke_test(self): ) def test_zeff_accepts_float_input(self, Z_eff: float): geo = circular_geometry.CircularConfig().build_geometry() - pc = plasma_composition.PlasmaComposition(Z_eff=Z_eff) # pyrefly: ignore[missing-argument] + pc = plasma_composition.PlasmaComposition(Z_eff=Z_eff) torax_pydantic.set_grid(pc, geo.torax_mesh) runtime_params = pc.build_runtime_params(t=0.0) # Check that the values in both Z_eff and Z_eff_face are the same @@ -74,7 +74,7 @@ def test_zeff_and_zeff_face_match_expected(self): } geo = circular_geometry.CircularConfig().build_geometry() - pc = plasma_composition.PlasmaComposition(Z_eff=zeff_profile) # pyrefly: ignore[missing-argument] + pc = plasma_composition.PlasmaComposition(Z_eff=zeff_profile) torax_pydantic.set_grid(pc, geo.torax_mesh) # Check values at t=0.0 @@ -129,7 +129,7 @@ def test_plasma_composition_under_jit(self, A_override): initial_zeff = 1.5 updated_zeff = 2.5 t = 0.0 - pc = plasma_composition.PlasmaComposition( # pyrefly: ignore[missing-argument] + pc = plasma_composition.PlasmaComposition( Z_eff=initial_zeff, A_i_override=A_override ) geo = circular_geometry.CircularConfig().build_geometry() diff --git a/torax/_src/core_profiles/tests/updaters_test.py b/torax/_src/core_profiles/tests/updaters_test.py index bf0823d71..0ecbb5891 100644 --- a/torax/_src/core_profiles/tests/updaters_test.py +++ b/torax/_src/core_profiles/tests/updaters_test.py @@ -41,13 +41,13 @@ def setUp(self): T_e = cell_variable.CellVariable( value=jnp.ones_like(self.geo.rho_norm), face_centers=self.geo.rho_face_norm, - right_face_constraint=1.0, # pyrefly: ignore[bad-argument-type] + right_face_constraint=1.0, right_face_grad_constraint=None, ) n_e = cell_variable.CellVariable( value=jnp.ones_like(self.geo.rho_norm), face_centers=self.geo.rho_face_norm, - right_face_constraint=1.0, # pyrefly: ignore[bad-argument-type] + right_face_constraint=1.0, right_face_grad_constraint=None, ) diff --git a/torax/_src/core_profiles/updaters.py b/torax/_src/core_profiles/updaters.py index 37368daf9..9435da0dc 100644 --- a/torax/_src/core_profiles/updaters.py +++ b/torax/_src/core_profiles/updaters.py @@ -288,7 +288,7 @@ def update_core_and_source_profiles_after_step( psidot = dataclasses.replace( core_profiles_t_plus_dt.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, ) diff --git a/torax/_src/edge/extended_lengyel/divertor_sol_1d.py b/torax/_src/edge/extended_lengyel/divertor_sol_1d.py index 034372ad7..54fefaea6 100644 --- a/torax/_src/edge/extended_lengyel/divertor_sol_1d.py +++ b/torax/_src/edge/extended_lengyel/divertor_sol_1d.py @@ -142,7 +142,7 @@ class DivertorSOL1D: state: ExtendedLengyelState @property - def electron_temp_at_cc_interface(self) -> jax.Array: + def electron_temp_at_cc_interface(self) -> array_typing.FloatScalar: """Calculates electron temperature at the convection/conduction interface. This function determines the electron temperature at the boundary between @@ -165,12 +165,12 @@ def electron_temp_at_cc_interface(self) -> jax.Array: self.state.T_e_target ) ) - return self.state.T_e_target / ( # pyrefly: ignore[bad-return] + return self.state.T_e_target / ( (1.0 - momentum_loss) / (2.0 * density_ratio) ) @property - def divertor_entrance_electron_temp(self) -> jax.Array: + def divertor_entrance_electron_temp(self) -> array_typing.FloatScalar: """Electron temperature at the divertor entrance [eV]. This formula is derived from the heat conduction equation integrated @@ -189,7 +189,7 @@ def divertor_entrance_electron_temp(self) -> jax.Array: ) ** (2.0 / 7.0) @property - def T_e_separatrix(self) -> jax.Array: + def T_e_separatrix(self) -> array_typing.FloatScalar: """Electron temperature at the separatrix [eV]. This formula is derived from the heat conduction equation integrated @@ -210,14 +210,14 @@ def T_e_separatrix(self) -> jax.Array: ) ** (2.0 / 7.0) @property - def separatrix_total_pressure(self) -> jax.Array: + def separatrix_total_pressure(self) -> array_typing.FloatScalar: """Total pressure at the separatrix [Pa]. This is the definition of total pressure (static + dynamic) at the separatrix, including both electron and ion contributions. """ return ( - (1.0 + self.params.mach_separatrix**2) # pyrefly: ignore[bad-return] + (1.0 + self.params.mach_separatrix**2) * self.params.separatrix_electron_density * self.T_e_separatrix * constants.CONSTANTS.eV_to_J @@ -229,7 +229,7 @@ def separatrix_total_pressure(self) -> jax.Array: ) @property - def required_power_loss(self) -> jax.Array: + def required_power_loss(self) -> array_typing.FloatScalar: """Required power loss fraction from the two-point model. Calculate momentum loss in the convection layer using an empirical fit. @@ -274,12 +274,12 @@ def required_power_loss(self) -> jax.Array: ) @property - def parallel_heat_flux_at_target(self) -> jax.Array: + def parallel_heat_flux_at_target(self) -> array_typing.FloatScalar: """Parallel heat flux at the divertor target [W/m^2].""" - return self.state.q_parallel * (1.0 - self.required_power_loss) # pyrefly: ignore[bad-return] + return self.state.q_parallel * (1.0 - self.required_power_loss) @property - def parallel_heat_flux_at_cc_interface(self) -> jax.Array: + def parallel_heat_flux_at_cc_interface(self) -> array_typing.FloatScalar: """Parallel heat flux at the convection-conduction interface [W/m^2]. Eq 29, Body et al. 2025. https://doi.org/10.1088/1741-4326/ade4d9 @@ -292,7 +292,7 @@ def parallel_heat_flux_at_cc_interface(self) -> jax.Array: return self.parallel_heat_flux_at_target / (1.0 - power_loss_conv_layer) @property - def divertor_Z_eff(self) -> jax.Array: + def divertor_Z_eff(self) -> array_typing.FloatScalar: return extended_lengyel_formulas.calc_Z_eff( c_z=self.state.c_z_prefactor, T_e=self.divertor_entrance_electron_temp / 1e3, # to keV @@ -303,7 +303,7 @@ def divertor_Z_eff(self) -> jax.Array: ) @property - def Z_eff_separatrix(self) -> jax.Array: + def Z_eff_separatrix(self) -> array_typing.FloatScalar: return extended_lengyel_formulas.calc_Z_eff( c_z=self.state.c_z_prefactor, T_e=self.T_e_separatrix / 1e3, # to keV @@ -328,7 +328,7 @@ def calc_q_parallel( params: ExtendedLengyelParameters, T_e_separatrix: array_typing.FloatScalar, alpha_t: array_typing.FloatScalar, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the parallel heat flux density. For the flux-tube assumed in the extended Lengyel model. @@ -392,7 +392,7 @@ def calc_q_parallel( * params.fieldline_pitch_at_omp ) - return q_parallel # pyrefly: ignore[bad-return] + return q_parallel def calc_alpha_t( @@ -475,7 +475,7 @@ def calc_alpha_t( def calc_T_e_target( sol_model: DivertorSOL1D, parallel_heat_flux_at_cc_interface: array_typing.FloatScalar, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculate the target electron temp from the two-point model. Args: @@ -552,7 +552,9 @@ def calc_T_e_target( return T_e_target_basic * f_vol_loss * f_other_T_e_target -def calc_kappa_e(Z_eff: array_typing.FloatScalar) -> jax.Array: +def calc_kappa_e( + Z_eff: array_typing.FloatScalar, +) -> array_typing.FloatScalar: """Corrected parallel electron heat conductivity prefactor. Eq 9, Body NF 2025. diff --git a/torax/_src/edge/extended_lengyel/extended_lengyel_formulas.py b/torax/_src/edge/extended_lengyel/extended_lengyel_formulas.py index 43dc974f7..6d3532d4f 100644 --- a/torax/_src/edge/extended_lengyel/extended_lengyel_formulas.py +++ b/torax/_src/edge/extended_lengyel/extended_lengyel_formulas.py @@ -122,7 +122,7 @@ def calc_separatrix_average_poloidal_field( plasma_current: array_typing.FloatScalar, minor_radius: array_typing.FloatScalar, shaping_factor: array_typing.FloatScalar, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the average poloidal field at the separatrix. Used for calculations related to magnetic geometry at the separatrix. @@ -142,7 +142,7 @@ def calc_separatrix_average_poloidal_field( The average poloidal field at the separatrix [T]. """ poloidal_circumference = 2.0 * jnp.pi * minor_radius * shaping_factor - return constants.CONSTANTS.mu_0 * plasma_current / poloidal_circumference # pyrefly: ignore[bad-return] + return constants.CONSTANTS.mu_0 * plasma_current / poloidal_circumference def calc_cylindrical_safety_factor( @@ -151,7 +151,7 @@ def calc_cylindrical_safety_factor( shaping_factor: array_typing.FloatScalar, minor_radius: array_typing.FloatScalar, major_radius: array_typing.FloatScalar, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the cylindrical safety factor. The cylindrical safety factor is a characteristic safety-factor value at the @@ -239,7 +239,7 @@ def calc_Z_eff( ne_tau: array_typing.FloatScalar, seed_impurity_weights: Mapping[str, array_typing.FloatScalar], fixed_impurity_concentrations: Mapping[str, array_typing.FloatScalar], -) -> jax.Array: +) -> array_typing.FloatScalar: """Helper function to calculate Z_eff in the extended Lengyel model. Z_eff is the effective ion charge, defined as sum(n_i * Z_i^2) / n_e. @@ -291,14 +291,14 @@ def calc_Z_eff( # Contribution from main ions n_i = (1 - dilution_factor) / Z_i Z_eff += n_i * Z_i**2 - return Z_eff[0] # Return scalar for extended-lengyel. # pyrefly: ignore[bad-index] + return jnp.squeeze(Z_eff) # Return scalar for extended-lengyel. def calc_enrichment_kallenbach( pressure_neutral_divertor: array_typing.FloatScalar, ion_symbol: str, enrichment_multiplier: array_typing.FloatScalar = 1.0, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculate divertor enrichment according to regression from Kallenbach 2024. A. Kallenbach et al 2024 Nucl. Fusion 64 056003 diff --git a/torax/_src/edge/extended_lengyel/extended_lengyel_model.py b/torax/_src/edge/extended_lengyel/extended_lengyel_model.py index f3518d7d5..2d24eb417 100644 --- a/torax/_src/edge/extended_lengyel/extended_lengyel_model.py +++ b/torax/_src/edge/extended_lengyel/extended_lengyel_model.py @@ -200,18 +200,18 @@ def __call__( # Calculate normalized poloidal flux (psi_norm) on the face grid. # Used to interpolate geometry quantities at psi_norm = 0.95. psi_face = core_profiles.psi.face_value() - psi_norm_face = (psi_face - psi_face[0]) / (psi_face[-1] - psi_face[0]) # pyrefly: ignore[bad-index] + psi_norm_face = (psi_face - psi_face[0]) / (psi_face[-1] - psi_face[0]) # Interpolate elongation and triangularity at psi_norm = 0.95 elongation_psi95 = jnp.interp(0.95, psi_norm_face, geo.elongation_face) triangularity_psi95 = jnp.interp(0.95, psi_norm_face, geo.delta_face) # Extract plasma state parameters from CoreProfiles at the LCFS - separatrix_electron_density = core_profiles.n_e.face_value()[-1] # pyrefly: ignore[bad-index] + separatrix_electron_density = core_profiles.n_e.face_value()[-1] # Calculate ion properties - n_i_sep = core_profiles.n_i.face_value()[-1] # pyrefly: ignore[bad-index] - n_imp_sep = core_profiles.n_impurity.face_value()[-1] # pyrefly: ignore[bad-index] + n_i_sep = core_profiles.n_i.face_value()[-1] + n_imp_sep = core_profiles.n_impurity.face_value()[-1] A_i_sep = core_profiles.A_i A_imp_sep = core_profiles.A_impurity_face[-1] mean_ion_charge_state = separatrix_electron_density / (n_i_sep + n_imp_sep) diff --git a/torax/_src/edge/extended_lengyel/extended_lengyel_solvers.py b/torax/_src/edge/extended_lengyel/extended_lengyel_solvers.py index 63784b429..2784a1359 100644 --- a/torax/_src/edge/extended_lengyel/extended_lengyel_solvers.py +++ b/torax/_src/edge/extended_lengyel/extended_lengyel_solvers.py @@ -286,7 +286,7 @@ def forward_mode_newton_solver( params = initial_sol_model.params residual_fun = functools.partial( - _forward_residual, params=params, fixed_cz=fixed_cz # pyrefly: ignore[bad-argument-type] + _forward_residual, params=params, fixed_cz=fixed_cz ) # 3. Run Newton-Raphson. @@ -367,7 +367,7 @@ def inverse_mode_newton_solver( params = initial_sol_model.params residual_fun = functools.partial( - _inverse_residual, params=params, fixed_Tt=fixed_Tt # pyrefly: ignore[bad-argument-type] + _inverse_residual, params=params, fixed_Tt=fixed_Tt ) # 3. Run Newton-Raphson. diff --git a/torax/_src/edge/extended_lengyel/pydantic_model.py b/torax/_src/edge/extended_lengyel/pydantic_model.py index ef67d6d84..e55921959 100644 --- a/torax/_src/edge/extended_lengyel/pydantic_model.py +++ b/torax/_src/edge/extended_lengyel/pydantic_model.py @@ -15,7 +15,7 @@ """Pydantic configs for all edge models, currently only extended_lengyel.""" import logging -from typing import Annotated, Any, Literal, Mapping +from typing import Annotated, Any, Literal, Mapping, Self import chex import jax.numpy as jnp import pydantic @@ -26,7 +26,6 @@ from torax._src.edge.extended_lengyel import extended_lengyel_formulas from torax._src.edge.extended_lengyel import extended_lengyel_model from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # pylint: disable=invalid-name @@ -223,7 +222,7 @@ def _set_default_fixed_point_iterations(cls, data: Any) -> Any: @pydantic.model_validator(mode='after') def _log_warning_for_unused_enrichment_factor( self, - ) -> typing_extensions.Self: + ) -> Self: """Logs a warning if enrichment_factor is provided when use_enrichment_model is True.""" if self.use_enrichment_model and self.enrichment_factor is not None: logging.warning( @@ -236,7 +235,7 @@ def _log_warning_for_unused_enrichment_factor( @pydantic.model_validator(mode='after') def _validate_enrichment_factor_keys( self, - ) -> typing_extensions.Self: + ) -> Self: """Validates that enrichment_factor keys are the same as impurity keys.""" if self.use_enrichment_model: @@ -280,7 +279,7 @@ def _validate_enrichment_factor_keys( @pydantic.model_validator(mode='after') def _validate_computation_mode_inputs( self, - ) -> typing_extensions.Self: + ) -> Self: """Validates inputs based on the specified computation mode.""" if self.computation_mode == extended_lengyel_enums.ComputationMode.FORWARD: if self.T_e_target is not None: diff --git a/torax/_src/edge/extended_lengyel/tests/extended_lengyel_model_test.py b/torax/_src/edge/extended_lengyel/tests/extended_lengyel_model_test.py index 26f5d1348..f8e139959 100644 --- a/torax/_src/edge/extended_lengyel/tests/extended_lengyel_model_test.py +++ b/torax/_src/edge/extended_lengyel/tests/extended_lengyel_model_test.py @@ -133,8 +133,8 @@ def test_call_inverse_mode(self): mock_geo ), qei=source_profiles.QeiInfo.zeros(mock_geo), - T_e={'generic_heat': el_heat}, # pyrefly: ignore[bad-argument-type, bad-assignment] - T_i={'generic_heat': ion_heat}, # pyrefly: ignore[bad-argument-type, bad-assignment] + T_e={'generic_heat': el_heat}, # pyrefly: ignore[bad-argument-type] + T_i={'generic_heat': ion_heat}, # pyrefly: ignore[bad-argument-type] ) # Verify that the mock sources integrate to the target power. diff --git a/torax/_src/edge/extended_lengyel/tests/extended_lengyel_multistart_test.py b/torax/_src/edge/extended_lengyel/tests/extended_lengyel_multistart_test.py index 60f513179..bcb055bc8 100644 --- a/torax/_src/edge/extended_lengyel/tests/extended_lengyel_multistart_test.py +++ b/torax/_src/edge/extended_lengyel/tests/extended_lengyel_multistart_test.py @@ -140,7 +140,7 @@ def solver_side_effect(initial_sol_model, **_): ) status = extended_lengyel_solvers.ExtendedLengyelSolverStatus( # pytype: disable=wrong-arg-types # pylint: disable=g-blanket-type-suppression - physics_outcome=phys_outcome, # pyrefly: ignore[bad-argument-type] + physics_outcome=phys_outcome, numerics_outcome=jax_root_finding.RootMetadata( iterations=jnp.array(5), error=error_val, diff --git a/torax/_src/edge/extended_lengyel/tests/extended_lengyel_output_test.py b/torax/_src/edge/extended_lengyel/tests/extended_lengyel_output_test.py index a89c2fdf3..aacd3a7bf 100644 --- a/torax/_src/edge/extended_lengyel/tests/extended_lengyel_output_test.py +++ b/torax/_src/edge/extended_lengyel/tests/extended_lengyel_output_test.py @@ -142,7 +142,7 @@ def test_roots_are_saved_correctly(self): Z_eff_separatrix=jnp.ones((num_roots,)) * 1.5, seed_impurity_concentrations={'Ne': jnp.ones((num_roots,)) * 0.01}, solver_status=extended_lengyel_solvers.ExtendedLengyelSolverStatus( - physics_outcome=jnp.array([ # pyrefly: ignore[bad-argument-type] + physics_outcome=jnp.array([ extended_lengyel_solvers.PhysicsOutcome.SUCCESS, extended_lengyel_solvers.PhysicsOutcome.SUCCESS, extended_lengyel_solvers.PhysicsOutcome.SUCCESS, diff --git a/torax/_src/fvm/calc_coeffs.py b/torax/_src/fvm/calc_coeffs.py index 6503e95f7..77557e7c8 100644 --- a/torax/_src/fvm/calc_coeffs.py +++ b/torax/_src/fvm/calc_coeffs.py @@ -34,7 +34,6 @@ from torax._src.sources import source_profile_builders from torax._src.sources import source_profiles as source_profiles_lib from torax._src.transport_model import transport_coefficients_builder -import typing_extensions # pylint: disable=invalid-name @@ -55,7 +54,9 @@ def __hash__(self) -> int: self.evolving_names, )) - def __eq__(self, other: typing_extensions.Self) -> bool: # pyrefly: ignore[bad-override] + def __eq__(self, other: object) -> bool: + if not isinstance(other, CoeffsCallback): + return False return ( self.models == other.models and self.evolving_names == other.evolving_names @@ -602,12 +603,12 @@ def _calc_coeffs_full( transient_in_cell=transient_in_cell, # pyrefly: ignore[bad-argument-type] d_face=d_face, # pyrefly: ignore[bad-argument-type] v_face=v_face, # pyrefly: ignore[bad-argument-type] - source_mat_cell=source_mat_cell, # pyrefly: ignore[bad-argument-type] - source_cell=source_cell, # pyrefly: ignore[bad-argument-type] - internal_boundary_condition_mask=internal_boundary_condition_mask, # pyrefly: ignore[bad-argument-type] + source_mat_cell=source_mat_cell, + source_cell=source_cell, + internal_boundary_condition_mask=internal_boundary_condition_mask, internal_boundary_condition_target_vec=( internal_boundary_condition_target_vec - ), # pyrefly: ignore[bad-argument-type] + ), ) return coeffs diff --git a/torax/_src/fvm/cell_variable.py b/torax/_src/fvm/cell_variable.py index e6a8f180d..fa6e1d209 100644 --- a/torax/_src/fvm/cell_variable.py +++ b/torax/_src/fvm/cell_variable.py @@ -168,10 +168,10 @@ class CellVariable: left_face_constraint: array_typing.FloatScalar | None = None right_face_constraint: array_typing.FloatScalar | None = None left_face_grad_constraint: array_typing.FloatScalar | None = ( - dataclasses.field(default_factory=_zero) # pyrefly: ignore[bad-assignment] + dataclasses.field(default_factory=_zero) ) right_face_grad_constraint: array_typing.FloatScalar | None = ( - dataclasses.field(default_factory=_zero) # pyrefly: ignore[bad-assignment] + dataclasses.field(default_factory=_zero) ) # Can't make the above default values be jax zeros because that would be a # call to jax before absl.app.run diff --git a/torax/_src/fvm/convection_terms.py b/torax/_src/fvm/convection_terms.py index 295381062..b0047e835 100644 --- a/torax/_src/fvm/convection_terms.py +++ b/torax/_src/fvm/convection_terms.py @@ -20,6 +20,7 @@ import chex import jax from jax import numpy as jnp +from torax._src import array_typing from torax._src import jax_utils from torax._src import tridiagonal from torax._src.fvm import cell_variable @@ -27,8 +28,8 @@ # TODO(b/469726859): Once non-uniform grid is supported add in testing. def make_convection_terms( - v_face: jax.Array, - d_face: jax.Array, + v_face: array_typing.Array, + d_face: array_typing.Array, var: cell_variable.CellVariable, dirichlet_mode: str = 'ghost', neumann_mode: str = 'ghost', @@ -66,6 +67,9 @@ def make_convection_terms( # Alpha weighting calculated using power law scheme described in # https://www.ctcms.nist.gov/fipy/documentation/numerical/scheme.html + v_face = jnp.asarray(v_face) + d_face = jnp.asarray(d_face) + # Avoid divide by zero eps = 1e-20 is_neg = d_face < 0.0 @@ -77,8 +81,8 @@ def make_convection_terms( ones = jnp.ones_like(v_face[1:-1]) scale = jnp.concatenate((half, ones, half)) - distance_to_left_ghost_cell_center = var.cell_widths[0] # pyrefly: ignore[bad-index] - distance_to_right_ghost_cell_center = var.cell_widths[-1] # pyrefly: ignore[bad-index] + distance_to_left_ghost_cell_center = var.cell_widths[0] + distance_to_right_ghost_cell_center = var.cell_widths[-1] cell_spacings = jnp.concat([ jnp.array([distance_to_left_ghost_cell_center]), var.cell_spacings, @@ -163,6 +167,7 @@ def peclet_to_alpha(p): raise ValueError(dirichlet_mode) else: # Gradient boundary condition at leftmost face + assert var.left_face_grad_constraint is not None diag_left_face = (v_face[0] - right_alpha[0] * v_face[1]) / cell_spacings[0] vec_left_face = ( -v_face[0] * (1.0 - left_alpha[0]) * var.left_face_grad_constraint @@ -207,6 +212,7 @@ def peclet_to_alpha(p): raise ValueError(dirichlet_mode) else: # Gradient boundary condition at rightmost face + assert var.right_face_grad_constraint is not None diag_right_face = ( -(v_face[-1] - v_face[-2] * left_alpha[-1]) / cell_spacings[-1] ) diff --git a/torax/_src/fvm/diffusion_terms.py b/torax/_src/fvm/diffusion_terms.py index 577164ab8..43832249d 100644 --- a/torax/_src/fvm/diffusion_terms.py +++ b/torax/_src/fvm/diffusion_terms.py @@ -43,8 +43,8 @@ def make_diffusion_terms( # Start by using the formula for the interior rows everywhere dx = var.cell_widths - distance_to_left_ghost_cell_center = dx[0] # pyrefly: ignore[bad-index] - distance_to_right_ghost_cell_center = dx[-1] # pyrefly: ignore[bad-index] + distance_to_left_ghost_cell_center = dx[0] + distance_to_right_ghost_cell_center = dx[-1] cell_spacings = jnp.concat([ jnp.array([distance_to_left_ghost_cell_center]), var.cell_spacings, @@ -53,12 +53,12 @@ def make_diffusion_terms( # Fill in the inner diagonal. face_flux_right = d_face[1:] / cell_spacings[1:] face_flux_left = d_face[:-1] / cell_spacings[:-1] - diag = (- face_flux_right - face_flux_left) / dx + diag = jnp.asarray((-face_flux_right - face_flux_left) / dx) off = d_face[1:-1] / var.cell_spacings # Divide by different cell widths for the upper and lower diagonals. - upper_off = off / dx[:-1] # pyrefly: ignore[bad-index] - lower_off = off / dx[1:] # pyrefly: ignore[bad-index] + upper_off = off / dx[:-1] + lower_off = off / dx[1:] vec = jnp.zeros_like(diag) @@ -83,21 +83,22 @@ def make_diffusion_terms( if var.left_face_constraint is not None: # Left face Dirichlet condition. - denom_left = cell_spacings[0] * dx[0] # pyrefly: ignore[bad-index] - denom_right = cell_spacings[1] * dx[0] # pyrefly: ignore[bad-index] - diag = diag.at[0].set(-2 * d_face[0] / denom_left - d_face[1] / denom_right) # pyrefly: ignore[missing-attribute] + denom_left = cell_spacings[0] * dx[0] + denom_right = cell_spacings[1] * dx[0] + diag = diag.at[0].set(-2 * d_face[0] / denom_left - d_face[1] / denom_right) vec = vec.at[0].set(2 * d_face[0] * var.left_face_constraint / denom_left) else: # Left face gradient condition. - denom_right = cell_spacings[1] * dx[0] # pyrefly: ignore[bad-index] - diag = diag.at[0].set(-d_face[1] / denom_right) # pyrefly: ignore[missing-attribute] + assert var.left_face_grad_constraint is not None + denom_right = cell_spacings[1] * dx[0] + diag = diag.at[0].set(-d_face[1] / denom_right) vec = vec.at[0].set( - -d_face[0] * var.left_face_grad_constraint / dx[0] # pyrefly: ignore[bad-index, unsupported-operation] + -d_face[0] * var.left_face_grad_constraint / dx[0] ) if var.right_face_constraint is not None: # Right face Dirichlet condition. - denom_left = cell_spacings[-2] * dx[-1] # pyrefly: ignore[bad-index] - denom_right = cell_spacings[-1] * dx[-1] # pyrefly: ignore[bad-index] + denom_left = cell_spacings[-2] * dx[-1] + denom_right = cell_spacings[-1] * dx[-1] diag = diag.at[-1].set( -2 * d_face[-1] / denom_right - d_face[-2] / denom_left ) @@ -106,10 +107,11 @@ def make_diffusion_terms( ) else: # Right face gradient condition. - denom_left = cell_spacings[-2] * dx[-1] # pyrefly: ignore[bad-index] + assert var.right_face_grad_constraint is not None + denom_left = cell_spacings[-2] * dx[-1] diag = diag.at[-1].set(-d_face[-2] / denom_left) vec = vec.at[-1].set( - d_face[-1] * var.right_face_grad_constraint / dx[-1] # pyrefly: ignore[bad-index, unsupported-operation] + d_face[-1] * var.right_face_grad_constraint / dx[-1] ) return tridiagonal.TriDiagonal(diag, upper_off, lower_off), vec diff --git a/torax/_src/fvm/tests/calc_coeffs_test.py b/torax/_src/fvm/tests/calc_coeffs_test.py index 3485b0e0b..d31291d85 100644 --- a/torax/_src/fvm/tests/calc_coeffs_test.py +++ b/torax/_src/fvm/tests/calc_coeffs_test.py @@ -232,7 +232,7 @@ def test_apply_transition_ramp_scaling_l_to_h(self): ), ) - scaled_pedestal_model_output = calc_coeffs._apply_transition_ramp_scaling( # pyrefly: ignore[bad-argument-type] + scaled_pedestal_model_output = calc_coeffs._apply_transition_ramp_scaling( pedestal_transition_state=state, ramp_fraction=0.5, ) diff --git a/torax/_src/geometry/base.py b/torax/_src/geometry/base.py index c05b1c9a4..1818e3ace 100644 --- a/torax/_src/geometry/base.py +++ b/torax/_src/geometry/base.py @@ -12,13 +12,12 @@ # See the License for the specific language governing permissions and # limitations under the License. """Base class for geometry configuration.""" -from typing import Annotated, Any +from typing import Annotated, Any, Self import numpy as np import pydantic from torax._src.torax_pydantic import interpolated_param_2d from torax._src.torax_pydantic import torax_pydantic -import typing_extensions class BaseGeometryConfig(torax_pydantic.BaseModelFrozen): @@ -52,7 +51,7 @@ def _validate_inputs(cls, data: Any) -> Any: return data @pydantic.model_validator(mode='after') - def _validate_n_rho_or_face_centers(self) -> typing_extensions.Self: + def _validate_n_rho_or_face_centers(self) -> Self: """Validates that there are at least 4 cells.""" if self.n_rho is None and self.face_centers is None: raise ValueError('Either n_rho or face_centers must be set.') @@ -82,7 +81,7 @@ def get_face_centers(self) -> np.ndarray: return self.face_centers return interpolated_param_2d.get_face_centers(self.n_rho) - def __eq__(self, other: typing_extensions.Self) -> bool: # pyrefly: ignore[bad-override] + def __eq__(self, other: object) -> bool: """Equality operator for BaseGeometryConfig.""" if not isinstance(other, type(self)): return False diff --git a/torax/_src/geometry/chease.py b/torax/_src/geometry/chease.py index 2fd24e529..18f431f3e 100644 --- a/torax/_src/geometry/chease.py +++ b/torax/_src/geometry/chease.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. """Functions for loading and representing a CHEASE geometry.""" -from typing import Annotated, Literal +from typing import Annotated, Literal, Self import numpy as np import pydantic from torax._src import constants @@ -21,7 +21,6 @@ from torax._src.geometry import geometry_loader from torax._src.geometry import standard_geometry from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # pylint: disable=invalid-name @@ -52,7 +51,7 @@ class CheaseConfig(base.BaseGeometryConfig): B_0: torax_pydantic.Tesla = 5.3 @pydantic.model_validator(mode='after') - def _check_fields(self) -> typing_extensions.Self: + def _check_fields(self) -> Self: if not self.R_major >= self.a_minor: raise ValueError('a_minor must be less than or equal to R_major.') return self diff --git a/torax/_src/geometry/circular_geometry.py b/torax/_src/geometry/circular_geometry.py index a1949f1c3..3cd8010b4 100644 --- a/torax/_src/geometry/circular_geometry.py +++ b/torax/_src/geometry/circular_geometry.py @@ -12,14 +12,12 @@ # See the License for the specific language governing permissions and # limitations under the License. """Classes for representing a circular geometry.""" -from typing import Annotated -from typing import Literal +from typing import Annotated, Literal, Self import numpy as np import pydantic from torax._src.geometry import base from torax._src.geometry import geometry from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # pylint: disable=invalid-name @@ -44,7 +42,7 @@ class CircularConfig(base.BaseGeometryConfig): elongation_LCFS: pydantic.PositiveFloat = 1.72 @pydantic.model_validator(mode='after') - def _check_fields(self) -> typing_extensions.Self: + def _check_fields(self) -> Self: if not self.R_major >= self.a_minor: raise ValueError('a_minor must be less than or equal to R_major.') return self diff --git a/torax/_src/geometry/eqdsk.py b/torax/_src/geometry/eqdsk.py index 23d4dcb41..cc7313d69 100644 --- a/torax/_src/geometry/eqdsk.py +++ b/torax/_src/geometry/eqdsk.py @@ -15,7 +15,7 @@ import json import logging -from typing import Annotated, Any, Literal +from typing import Annotated, Any, Literal, Self import contourpy import eqdsk @@ -31,7 +31,6 @@ from torax._src.geometry import geometry_loader from torax._src.geometry import standard_geometry from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # Inject `npt` into eqdsk.file's runtime namespace to prevent Pydantic from # raising `PydanticUndefinedAnnotation: name 'npt' is not defined.` @@ -105,7 +104,7 @@ def _eqdskinterface_serializer( return json.loads(json_str) @pydantic.model_validator(mode='after') - def _validate_model(self) -> typing_extensions.Self: + def _validate_model(self) -> Self: if self.geometry_file is None and self.eqdsk_object is None: raise ValueError( "Either 'geometry_file' or 'eqdsk_object' must be provided." diff --git a/torax/_src/geometry/fbt.py b/torax/_src/geometry/fbt.py index f6f9471ab..972b8bdc9 100644 --- a/torax/_src/geometry/fbt.py +++ b/torax/_src/geometry/fbt.py @@ -16,9 +16,7 @@ from collections.abc import Mapping import enum import logging -from typing import Annotated -from typing import Any -from typing import Literal, TypeAlias +from typing import Annotated, Any, Literal, Self, TypeAlias import jax import numpy as np @@ -31,7 +29,6 @@ from torax._src.geometry import geometry_provider from torax._src.geometry import standard_geometry from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # pylint: disable=invalid-name LY_OBJECT_TYPE: TypeAlias = ( @@ -112,7 +109,7 @@ def _conform_data(cls, data: dict[str, Any]) -> dict[str, Any]: return data @pydantic.model_validator(mode='after') - def _validate_model(self) -> typing_extensions.Self: + def _validate_model(self) -> Self: if self.LY_bundle_object is not None and self.LY_object is not None: raise ValueError( "Cannot use 'LY_object' together with a bundled FBT file" diff --git a/torax/_src/geometry/geometry.py b/torax/_src/geometry/geometry.py index bb7092df5..c73fe7f6c 100644 --- a/torax/_src/geometry/geometry.py +++ b/torax/_src/geometry/geometry.py @@ -230,7 +230,9 @@ class Geometry: Phi_b_dot: array_typing.FloatScalar _z_magnetic_axis: array_typing.FloatScalar | None - def __eq__(self, other: 'Geometry') -> bool: # pyrefly: ignore[bad-override] + def __eq__(self, other: object) -> bool: + if not isinstance(other, Geometry): + return False try: chex.assert_trees_all_equal(self, other) except AssertionError: @@ -358,12 +360,12 @@ def g1_over_vpr2_face(self) -> jax.Array: ) @property - def gm9(self) -> jax.Array: + def gm9(self) -> array_typing.Array: r"""<1/R> on cell grid [:math:`\mathrm{m}^{-1}`].""" - return 2 * jnp.pi * self.spr / self.vpr # pyrefly: ignore[bad-return] + return 2 * jnp.pi * self.spr / self.vpr @property - def gm9_face(self) -> jax.Array: + def gm9_face(self) -> array_typing.Array: r"""<1/R> on face grid [:math:`\mathrm{m}^{-1}`].""" bulk = 2 * jnp.pi * self.spr_face[..., 1:] / self.vpr_face[..., 1:] first_element = 1 / self.R_major_profile_face[..., 0] @@ -372,16 +374,16 @@ def gm9_face(self) -> jax.Array: ) @property - def R_major_profile(self) -> jax.Array: + def R_major_profile(self) -> array_typing.Array: """Local major radius on cell grid [m].""" - return (self.R_in + self.R_out) / 2 # pyrefly: ignore[bad-return] + return (self.R_in + self.R_out) / 2 @property - def R_major_profile_face(self) -> jax.Array: + def R_major_profile_face(self) -> array_typing.Array: """Local major radius on face grid [m].""" - return (self.R_in_face + self.R_out_face) / 2 # pyrefly: ignore[bad-return] + return (self.R_in_face + self.R_out_face) / 2 - def z_magnetic_axis(self) -> chex.Numeric: + def z_magnetic_axis(self) -> array_typing.FloatScalar: """z position of magnetic axis [m].""" z_magnetic_axis = self._z_magnetic_axis if z_magnetic_axis is not None: diff --git a/torax/_src/geometry/geometry_provider.py b/torax/_src/geometry/geometry_provider.py index 94b583d03..57f76d50e 100644 --- a/torax/_src/geometry/geometry_provider.py +++ b/torax/_src/geometry/geometry_provider.py @@ -20,7 +20,7 @@ from collections.abc import Mapping import dataclasses import functools -from typing import Protocol, Type +from typing import Protocol, Self, Type import chex import jax @@ -30,7 +30,6 @@ from torax._src import jax_utils from torax._src.geometry import geometry from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # Using invalid-name because we are using the same naming convention as the # external physics implementations @@ -165,7 +164,7 @@ def create_provider( cls, geometries: Mapping[float, geometry.Geometry], calcphibdot: bool, - ) -> typing_extensions.Self: + ) -> Self: """Creates a GeometryProvider from a mapping of times to geometries.""" # Create a list of times and geometries. times = np.asarray(list(geometries.keys()), dtype=jax_utils.get_np_dtype()) diff --git a/torax/_src/geometry/imas.py b/torax/_src/geometry/imas.py index d0296f260..7c19ae5c9 100644 --- a/torax/_src/geometry/imas.py +++ b/torax/_src/geometry/imas.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. """Functions for loading and representing an IMAS geometry.""" -from typing import Annotated, Literal +from typing import Annotated, Literal, Self from imas import ids_toplevel import pydantic @@ -22,7 +22,6 @@ from torax._src.geometry import standard_geometry from torax._src.imas_tools.input import equilibrium as imas_geometry from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # pylint: disable=invalid-name @@ -107,7 +106,7 @@ class IMASConfig(base.BaseGeometryConfig): ] = False @pydantic.model_validator(mode='after') - def _validate_model(self) -> typing_extensions.Self: + def _validate_model(self) -> Self: specified_inputs = [ field for field in [ diff --git a/torax/_src/imas_tools/output/core_profiles.py b/torax/_src/imas_tools/output/core_profiles.py index fb452ff81..4d57e2e8b 100644 --- a/torax/_src/imas_tools/output/core_profiles.py +++ b/torax/_src/imas_tools/output/core_profiles.py @@ -294,11 +294,11 @@ def _fill_profiles_1d_grid( [[0.0], geometry_slice.rho, [geometry_slice.rho_b]] ) ids.profiles_1d[i].grid.psi = cp_state.psi.cell_plus_boundaries() - ids.profiles_1d[i].grid.psi_magnetic_axis = cp_state.psi.left_face_value[0] # pyrefly: ignore[bad-index] - ids.profiles_1d[i].grid.psi_boundary = cp_state.psi.right_face_value[0] # pyrefly: ignore[bad-index] + ids.profiles_1d[i].grid.psi_magnetic_axis = cp_state.psi.left_face_value[0] + ids.profiles_1d[i].grid.psi_boundary = cp_state.psi.right_face_value[0] ids.profiles_1d[i].grid.rho_pol_norm = np.sqrt( - (cp_state.psi.cell_plus_boundaries() - cp_state.psi.left_face_value[0]) # pyrefly: ignore[bad-index] - / (cp_state.psi.right_face_value[0] - cp_state.psi.left_face_value[0]) # pyrefly: ignore[bad-index] + (cp_state.psi.cell_plus_boundaries() - cp_state.psi.left_face_value[0]) + / (cp_state.psi.right_face_value[0] - cp_state.psi.left_face_value[0]) ) ids.profiles_1d[i].grid.volume = output.extend_cell_grid_to_boundaries( [geometry_slice.volume], np.array([geometry_slice.volume_face]) # pyrefly: ignore[bad-argument-type] diff --git a/torax/_src/imas_tools/output/equilibrium.py b/torax/_src/imas_tools/output/equilibrium.py index f6186c6a4..712ba388e 100644 --- a/torax/_src/imas_tools/output/equilibrium.py +++ b/torax/_src/imas_tools/output/equilibrium.py @@ -68,8 +68,8 @@ def torax_state_to_imas_equilibrium( eq.boundary.geometric_axis.r = geometry.R_major eq.boundary.minor_radius = geometry.a_minor eq.profiles_1d.psi = core_profiles.psi.face_value() - psi_axis = core_profiles.psi.face_value()[0] # pyrefly: ignore[bad-index] - psi_boundary = core_profiles.psi.face_value()[-1] # pyrefly: ignore[bad-index] + psi_axis = core_profiles.psi.face_value()[0] + psi_boundary = core_profiles.psi.face_value()[-1] eq.global_quantities.psi_axis = psi_axis eq.global_quantities.psi_boundary = psi_boundary eq.profiles_1d.psi_norm = (core_profiles.psi.face_value() - psi_axis) / ( diff --git a/torax/_src/interpolated_param.py b/torax/_src/interpolated_param.py index fc33f2647..cfa809a75 100644 --- a/torax/_src/interpolated_param.py +++ b/torax/_src/interpolated_param.py @@ -473,7 +473,9 @@ def param(self) -> InterpolatedParamBase: """Returns the JAX-friendly interpolated param used under the hood.""" return self._param - def __eq__(self, other: 'InterpolatedVarSingleAxis') -> bool: # pyrefly: ignore[bad-override] + def __eq__(self, other: object) -> bool: + if not isinstance(other, InterpolatedVarSingleAxis): + return False try: chex.assert_trees_all_equal(self, other) except AssertionError: diff --git a/torax/_src/math_utils.py b/torax/_src/math_utils.py index d2f448017..990f4161d 100644 --- a/torax/_src/math_utils.py +++ b/torax/_src/math_utils.py @@ -224,7 +224,7 @@ def cell_integration( @array_typing.jaxtyped def area_integration( - value: array_typing.FloatVector, + value: array_typing.FloatVectorCell, geo: geometry.Geometry, ) -> array_typing.FloatScalar: """Calculates integral of value using an area metric.""" @@ -233,7 +233,7 @@ def area_integration( @array_typing.jaxtyped def volume_integration( - value: array_typing.FloatVector, + value: array_typing.FloatVectorCell, geo: geometry.Geometry, ) -> array_typing.FloatScalar: """Calculates integral of value using a volume metric.""" @@ -242,7 +242,7 @@ def volume_integration( @array_typing.jaxtyped def line_average( - value: array_typing.FloatVector, + value: array_typing.FloatVectorCell, geo: geometry.Geometry, ) -> array_typing.FloatScalar: """Calculates line-averaged value from input profile.""" @@ -251,7 +251,7 @@ def line_average( @array_typing.jaxtyped def volume_average( - value: array_typing.FloatVector, + value: array_typing.FloatVectorCell, geo: geometry.Geometry, ) -> array_typing.FloatScalar: """Calculates volume-averaged value from input profile.""" @@ -300,9 +300,9 @@ def cumulative_volume_integration( return cumulative_cell_integration(value * geo.vpr, geo) -def safe_divide( - *, num: chex.Array, denom: chex.Array, eps: float -) -> chex.Array: +def safe_divide[T: chex.Numeric]( + *, num: T, denom: chex.Numeric, eps: chex.Numeric +) -> T: """Divides y by x, adding eps to the denominator for numerical stability. Args: diff --git a/torax/_src/mhd/sawtooth/flatten_profile.py b/torax/_src/mhd/sawtooth/flatten_profile.py index 6d6e67388..38e30d592 100644 --- a/torax/_src/mhd/sawtooth/flatten_profile.py +++ b/torax/_src/mhd/sawtooth/flatten_profile.py @@ -207,8 +207,8 @@ def flatten_current_profile( # face boundary. new_psi = ( new_psi.value - - new_psi.face_value()[-1] # pyrefly: ignore[bad-index] - + original_psi_profile.face_value()[-1] # pyrefly: ignore[bad-index] + - new_psi.face_value()[-1] + + original_psi_profile.face_value()[-1] ) return dataclasses.replace( diff --git a/torax/_src/mhd/sawtooth/tests/flatten_profile_test.py b/torax/_src/mhd/sawtooth/tests/flatten_profile_test.py index ffe56000b..58b1f3b7a 100644 --- a/torax/_src/mhd/sawtooth/tests/flatten_profile_test.py +++ b/torax/_src/mhd/sawtooth/tests/flatten_profile_test.py @@ -197,12 +197,12 @@ def test_flatten_profile_logic_and_conservation( with self.subTest('conservation_within_mixing_radius'): self._check_conservation_within_mixing_radius( - initial_profile.value, flattened_profile.value, rho_norm_mixing # pyrefly: ignore[bad-argument-type] + initial_profile.value, flattened_profile.value, rho_norm_mixing ) with self.subTest('total_conservation'): self._check_total_conservation( - initial_profile.value, flattened_profile.value # pyrefly: ignore[bad-argument-type] + initial_profile.value, flattened_profile.value ) # Detailed checks on profile shape @@ -213,8 +213,8 @@ def test_flatten_profile_logic_and_conservation( with self.subTest('outer_region_unchanged'): if idx_mixing < _NRHO: np.testing.assert_allclose( - val_after[idx_mixing:], # pyrefly: ignore[bad-index] - initial_profile.value[idx_mixing:], # pyrefly: ignore[bad-index] + val_after[idx_mixing:], + initial_profile.value[idx_mixing:], err_msg='Profile changed outside mixing radius', ) @@ -284,22 +284,22 @@ def test_temperature_profile_flattening_and_energy_conservation( ) initial_pressure_profile = self._create_profile( - initial_temperature_profile.value * initial_density_profile.value # pyrefly: ignore[bad-argument-type] + initial_temperature_profile.value * initial_density_profile.value ) flattened_pressure_profile = self._create_profile( - flattened_temperature_profile.value * flattened_density_profile.value # pyrefly: ignore[bad-argument-type] + flattened_temperature_profile.value * flattened_density_profile.value ) with self.subTest('conservation_within_mixing_radius'): self._check_conservation_within_mixing_radius( - initial_pressure_profile.value, # pyrefly: ignore[bad-argument-type] - flattened_pressure_profile.value, # pyrefly: ignore[bad-argument-type] + initial_pressure_profile.value, + flattened_pressure_profile.value, rho_norm_mixing, ) with self.subTest('total_conservation'): self._check_total_conservation( - initial_pressure_profile.value, flattened_pressure_profile.value # pyrefly: ignore[bad-argument-type] + initial_pressure_profile.value, flattened_pressure_profile.value ) # pylint: disable=invalid-name diff --git a/torax/_src/neoclassical/bootstrap_current/redl.py b/torax/_src/neoclassical/bootstrap_current/redl.py index 7b7104a9f..b979bb5a2 100644 --- a/torax/_src/neoclassical/bootstrap_current/redl.py +++ b/torax/_src/neoclassical/bootstrap_current/redl.py @@ -129,16 +129,16 @@ def _calculate_bootstrap_current( nu_e_star = formulas.calculate_nu_e_star( q=q_face, geo=geo, - n_e=n_e.face_value(), # pyrefly: ignore[bad-argument-type] - T_e=T_e.face_value(), # pyrefly: ignore[bad-argument-type] + n_e=n_e.face_value(), + T_e=T_e.face_value(), Z_eff=Z_eff_face, log_lambda_ei=log_lambda_ei, ) nu_i_star = formulas.calculate_nu_i_star( q=q_face, geo=geo, - n_i=n_i.face_value(), # pyrefly: ignore[bad-argument-type] - T_i=T_i.face_value(), # pyrefly: ignore[bad-argument-type] + n_i=n_i.face_value(), + T_i=T_i.face_value(), Z_eff=Z_eff_face, log_lambda_ii=log_lambda_ii, ) diff --git a/torax/_src/neoclassical/bootstrap_current/sauter.py b/torax/_src/neoclassical/bootstrap_current/sauter.py index 4cac40b1c..160efdfc2 100644 --- a/torax/_src/neoclassical/bootstrap_current/sauter.py +++ b/torax/_src/neoclassical/bootstrap_current/sauter.py @@ -120,16 +120,16 @@ def _calculate_bootstrap_current( nu_e_star = formulas.calculate_nu_e_star( q=q_face, geo=geo, - n_e=n_e.face_value(), # pyrefly: ignore[bad-argument-type] - T_e=T_e.face_value(), # pyrefly: ignore[bad-argument-type] + n_e=n_e.face_value(), + T_e=T_e.face_value(), Z_eff=Z_eff_face, log_lambda_ei=log_lambda_ei, ) nu_i_star = formulas.calculate_nu_i_star( q=q_face, geo=geo, - n_i=n_i.face_value(), # pyrefly: ignore[bad-argument-type] - T_i=T_i.face_value(), # pyrefly: ignore[bad-argument-type] + n_i=n_i.face_value(), + T_i=T_i.face_value(), Z_eff=Z_eff_face, log_lambda_ii=log_lambda_ii, ) diff --git a/torax/_src/neoclassical/conductivity/sauter.py b/torax/_src/neoclassical/conductivity/sauter.py index 1c50e114e..57ec2c817 100644 --- a/torax/_src/neoclassical/conductivity/sauter.py +++ b/torax/_src/neoclassical/conductivity/sauter.py @@ -110,8 +110,8 @@ def _calculate_conductivity( nu_e_star_face = formulas.calculate_nu_e_star( q=q_face, geo=geo, - n_e=n_e.face_value(), # pyrefly: ignore[bad-argument-type] - T_e=T_e.face_value(), # pyrefly: ignore[bad-argument-type] + n_e=n_e.face_value(), + T_e=T_e.face_value(), Z_eff=Z_eff_face, log_lambda_ei=log_lambda_ei, ) diff --git a/torax/_src/neoclassical/formulas/formulas.py b/torax/_src/neoclassical/formulas/formulas.py index 490d553ad..3a7430b4b 100644 --- a/torax/_src/neoclassical/formulas/formulas.py +++ b/torax/_src/neoclassical/formulas/formulas.py @@ -231,7 +231,7 @@ def calculate_poloidal_velocity( q=q, geo=geo, n_i=n_i, - T_i=T_i_face, # pyrefly: ignore[bad-argument-type] + T_i=T_i_face, Z_eff=Z_eff, log_lambda_ii=log_lambda_ii, ) @@ -314,7 +314,7 @@ def calculate_analytic_bootstrap_current( dlnte_drnorm = T_e.face_grad() / T_e.face_value() dlnti_drnorm = T_i.face_grad() / T_i.face_value() - global_coeff = prefactor[1:] / dpsi_drnorm[1:] # pyrefly: ignore[bad-index] + global_coeff = prefactor[1:] / dpsi_drnorm[1:] global_coeff = jnp.concatenate([jnp.zeros(1), global_coeff]) necoeff = L31 * pe diff --git a/torax/_src/neoclassical/formulas/tests/formulas_test.py b/torax/_src/neoclassical/formulas/tests/formulas_test.py index bf4c329fe..a8faf9042 100644 --- a/torax/_src/neoclassical/formulas/tests/formulas_test.py +++ b/torax/_src/neoclassical/formulas/tests/formulas_test.py @@ -87,8 +87,8 @@ def setUp(self): self.nu_e_star = formulas.calculate_nu_e_star( q=self.core_profiles.q_face, geo=self.geo, - n_e=self.core_profiles.n_e.face_value(), # pyrefly: ignore[bad-argument-type] - T_e=self.core_profiles.T_e.face_value(), # pyrefly: ignore[bad-argument-type] + n_e=self.core_profiles.n_e.face_value(), + T_e=self.core_profiles.T_e.face_value(), Z_eff=self.core_profiles.Z_eff_face, log_lambda_ei=log_lambda_ei, ) diff --git a/torax/_src/neoclassical/formulas/tests/redl_test.py b/torax/_src/neoclassical/formulas/tests/redl_test.py index 338206ed7..d5164ed25 100644 --- a/torax/_src/neoclassical/formulas/tests/redl_test.py +++ b/torax/_src/neoclassical/formulas/tests/redl_test.py @@ -82,8 +82,8 @@ def setUp(self): self.nu_e_star = formulas.calculate_nu_e_star( q=self.core_profiles.q_face, geo=self.geo, - n_e=self.core_profiles.n_e.face_value(), # pyrefly: ignore[bad-argument-type] - T_e=self.core_profiles.T_e.face_value(), # pyrefly: ignore[bad-argument-type] + n_e=self.core_profiles.n_e.face_value(), + T_e=self.core_profiles.T_e.face_value(), Z_eff=self.core_profiles.Z_eff_face, log_lambda_ei=log_lambda_ei, ) diff --git a/torax/_src/neoclassical/formulas/tests/sauter_test.py b/torax/_src/neoclassical/formulas/tests/sauter_test.py index dd0816f5b..a5d23c7aa 100644 --- a/torax/_src/neoclassical/formulas/tests/sauter_test.py +++ b/torax/_src/neoclassical/formulas/tests/sauter_test.py @@ -85,8 +85,8 @@ def setUp(self): self.nu_e_star = formulas.calculate_nu_e_star( q=self.core_profiles.q_face, geo=self.geo, - n_e=self.core_profiles.n_e.face_value(), # pyrefly: ignore[bad-argument-type] - T_e=self.core_profiles.T_e.face_value(), # pyrefly: ignore[bad-argument-type] + n_e=self.core_profiles.n_e.face_value(), + T_e=self.core_profiles.T_e.face_value(), Z_eff=self.core_profiles.Z_eff_face, log_lambda_ei=log_lambda_ei, ) diff --git a/torax/_src/neoclassical/transport/angioni_sauter.py b/torax/_src/neoclassical/transport/angioni_sauter.py index ea66d963b..84406dc48 100644 --- a/torax/_src/neoclassical/transport/angioni_sauter.py +++ b/torax/_src/neoclassical/transport/angioni_sauter.py @@ -22,7 +22,7 @@ """ import dataclasses -from typing import Annotated, Literal +from typing import Annotated, Literal, override import jax from jax import numpy as jnp @@ -39,7 +39,6 @@ from torax._src.neoclassical.transport import runtime_params as transport_runtime_params from torax._src.physics import collisions from torax._src.torax_pydantic import torax_pydantic -from typing_extensions import override # pylint: disable=invalid-name @@ -196,8 +195,8 @@ def _calculate_angioni_sauter_transport( nu_e_star = formulas.calculate_nu_e_star( q=core_profiles.q_face, geo=geometry, - n_e=core_profiles.n_e.face_value(), # pyrefly: ignore[bad-argument-type] - T_e=core_profiles.T_e.face_value(), # pyrefly: ignore[bad-argument-type] + n_e=core_profiles.n_e.face_value(), + T_e=core_profiles.T_e.face_value(), Z_eff=core_profiles.Z_eff_face, log_lambda_ei=log_lambda_ei, ) @@ -311,17 +310,17 @@ def _calculate_angioni_sauter_transport( # to avoid division by near-zero and unphysical values. chi_neo_e_bulk = -Be2[1:] / ( - core_profiles.n_e.face_value()[1:] # pyrefly: ignore[bad-index] + core_profiles.n_e.face_value()[1:] * dlnte_dpsi[1:] # pyrefly: ignore[bad-index] - * (dpsi_drhon[1:] / geometry.rho_b) ** 2 # pyrefly: ignore[bad-index] + * (dpsi_drhon[1:] / geometry.rho_b) ** 2 + constants.CONSTANTS.eps ) chi_neo_e = jnp.concatenate([chi_neo_e_bulk[0:1], chi_neo_e_bulk]) chi_neo_i_bulk = -Bi2[1:] / ( - core_profiles.n_i.face_value()[1:] # pyrefly: ignore[bad-index] + core_profiles.n_i.face_value()[1:] * dlnti_dpsi[1:] # pyrefly: ignore[bad-index] - * (dpsi_drhon[1:] / geometry.rho_b) ** 2 # pyrefly: ignore[bad-index] + * (dpsi_drhon[1:] / geometry.rho_b) ** 2 + constants.CONSTANTS.eps ) chi_neo_i = jnp.concatenate([chi_neo_i_bulk[0:1], chi_neo_i_bulk]) @@ -331,8 +330,8 @@ def _calculate_angioni_sauter_transport( # Diffusive part of particle flux # D_e * dn_e/drho = - L00 *dlog(n_e)/dpsi / dpsi/drho D_neo_e_bulk = -Lmn_e[1:, 0, 0] / ( - core_profiles.n_e.face_value()[1:] # pyrefly: ignore[bad-index] - * (dpsi_drhon[1:] / geometry.rho_b) ** 2 # pyrefly: ignore[bad-index] + core_profiles.n_e.face_value()[1:] + * (dpsi_drhon[1:] / geometry.rho_b) ** 2 + constants.CONSTANTS.eps ) D_neo_e = jnp.concatenate([D_neo_e_bulk[0:1], D_neo_e_bulk]) @@ -343,12 +342,12 @@ def _calculate_angioni_sauter_transport( V_neo_e_bulk = ( (Lmn_e[1:, 0, 0] + Lmn_e[1:, 0, 1]) * dlnte_dpsi[1:] # pyrefly: ignore[bad-index] + (1 - Rpe[1:]) / Rpe[1:] * Lmn_e[1:, 0, 0] * dlnni_dpsi[1:] # pyrefly: ignore[bad-index] - + (1 - Rpe[1:]) # pyrefly: ignore[bad-index] - / Rpe[1:] # pyrefly: ignore[bad-index] + + (1 - Rpe[1:]) + / Rpe[1:] * (Lmn_e[1:, 0, 0] + alpha[1:] * Lmn_e[1:, 0, 3]) * dlnti_dpsi[1:] # pyrefly: ignore[bad-index] ) / ( - dpsi_drhon[1:] / geometry.rho_b * core_profiles.n_e.face_value()[1:] # pyrefly: ignore[bad-index] + dpsi_drhon[1:] / geometry.rho_b * core_profiles.n_e.face_value()[1:] + constants.CONSTANTS.eps ) V_neo_e = jnp.concatenate([V_neo_e_bulk[0:1], V_neo_e_bulk]) @@ -360,8 +359,8 @@ def _calculate_angioni_sauter_transport( * E_parallel[1:] / ( geometry.B_0 - * (dpsi_drhon[1:] / geometry.rho_b) # pyrefly: ignore[bad-index] - * core_profiles.n_e.face_value()[1:] # pyrefly: ignore[bad-index] + * (dpsi_drhon[1:] / geometry.rho_b) + * core_profiles.n_e.face_value()[1:] + constants.CONSTANTS.eps ) ) @@ -744,7 +743,7 @@ def _calculate_shaing_transport( # (currently we simply copy the value at i=1). This is ok as chi[0] is never # used. dpsi_drhon = core_profiles.psi.face_grad() - dpsi_drhon = dpsi_drhon.at[0].set(dpsi_drhon[1]) # pyrefly: ignore[bad-index, missing-attribute] + dpsi_drhon = dpsi_drhon.at[0].set(dpsi_drhon[1]) # pyrefly: ignore[missing-attribute] conversion_factor = 1 / (dpsi_drhon / (2 * jnp.pi * geometry.rho_b)) ** 2 # Trapped particle fraction (Equation 46, Shaing March 1997) diff --git a/torax/_src/neoclassical/transport/base.py b/torax/_src/neoclassical/transport/base.py index 0a72d9e08..e5ef31ba6 100644 --- a/torax/_src/neoclassical/transport/base.py +++ b/torax/_src/neoclassical/transport/base.py @@ -15,6 +15,7 @@ """Base class for neoclassical transport models.""" import abc import dataclasses +from typing import Self import jax import jax.numpy as jnp @@ -25,7 +26,6 @@ from torax._src.geometry import geometry as geometry_lib from torax._src.neoclassical.transport import runtime_params as transport_runtime_params from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # pylint: disable=invalid-name @@ -157,7 +157,7 @@ class NeoclassicalTransportModelConfig(torax_pydantic.BaseModelFrozen, abc.ABC): V_e_max: torax_pydantic.MeterPerSecond = 50.0 @pydantic.model_validator(mode='after') - def _check_fields(self) -> typing_extensions.Self: + def _check_fields(self) -> Self: if not self.chi_min < self.chi_max: raise ValueError('chi_min must be less than chi_max.') if not self.D_e_min < self.D_e_max: diff --git a/torax/_src/neoclassical/transport/zeros.py b/torax/_src/neoclassical/transport/zeros.py index a3231d162..c2b5493ad 100644 --- a/torax/_src/neoclassical/transport/zeros.py +++ b/torax/_src/neoclassical/transport/zeros.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. """Zeros model for neoclassical transport.""" -from typing import Annotated, Literal +from typing import Annotated, Literal, override import jax.numpy as jnp from torax._src import state @@ -20,7 +20,6 @@ from torax._src.geometry import geometry as geometry_lib from torax._src.neoclassical.transport import base from torax._src.torax_pydantic import torax_pydantic -from typing_extensions import override class ZerosModel(base.NeoclassicalTransportModel): diff --git a/torax/_src/orchestration/step_function.py b/torax/_src/orchestration/step_function.py index dd9cf39f6..970781cfb 100644 --- a/torax/_src/orchestration/step_function.py +++ b/torax/_src/orchestration/step_function.py @@ -343,11 +343,11 @@ def body(args): # Set the dt to the original dt passed to the function, and the t to the # final time. # In case we exited early return the actual elapsed dt. - elapsed_dt = dt - remaining_dt # pyrefly: ignore[unsupported-operation] + elapsed_dt = dt - remaining_dt output_state = dataclasses.replace( output_state, t=input_state.t + elapsed_dt, - dt=elapsed_dt, # pyrefly: ignore[bad-argument-type] + dt=elapsed_dt, ) return output_state, post_processed_outputs diff --git a/torax/_src/orchestration/step_function_processing.py b/torax/_src/orchestration/step_function_processing.py index 7445589c8..da489d929 100644 --- a/torax/_src/orchestration/step_function_processing.py +++ b/torax/_src/orchestration/step_function_processing.py @@ -17,6 +17,7 @@ import dataclasses import jax import jax.numpy as jnp +from torax._src import array_typing from torax._src import models as models_lib from torax._src import state from torax._src.config import build_runtime_params @@ -81,7 +82,7 @@ def _update_pedestal_transition_state( # Calculate P_SOL (total power crossing the separatrix). P_SOL = power_scaling_formation_model_lib.calculate_P_SOL_total( - internal_plasma_energy=core_profiles.internal_plasma_energy, # pyrefly: ignore[bad-argument-type] + internal_plasma_energy=core_profiles.internal_plasma_energy, core_sources=core_sources, geo=geo, include_dW_dt=runtime_params.pedestal.include_dW_dt_in_P_SOL, @@ -125,8 +126,8 @@ def _update_adaptive_transport( pedestal_transition_state_lib.PedestalTransitionState ), runtime_params: runtime_params_lib.RuntimeParams, - P_SOL: jax.Array, - P_LH: jax.Array, + P_SOL: array_typing.FloatScalar, + P_LH: array_typing.FloatScalar, ) -> pedestal_transition_state_lib.PedestalTransitionState: """Updates pedestal transition state for ADAPTIVE_TRANSPORT mode. @@ -181,8 +182,8 @@ def _update_internal_boundary_condition( core_profiles: state.CoreProfiles, core_sources: source_profiles_lib.SourceProfiles, models: models_lib.Models, - P_SOL: jax.Array, - P_LH: jax.Array, + P_SOL: array_typing.FloatScalar, + P_LH: array_typing.FloatScalar, ) -> pedestal_transition_state_lib.PedestalTransitionState: """Updates pedestal transition state for INTERNAL_BOUNDARY_CONDITION mode. @@ -322,17 +323,17 @@ def _update_internal_boundary_condition( ) new_T_i_ped_L_mode = jnp.where( update_L_mode_values, - core_profiles.T_i.value[ped_top_idx], # pyrefly: ignore[bad-index] + core_profiles.T_i.value[ped_top_idx], pedestal_transition_state.T_i_ped_L_mode, ) new_T_e_ped_L_mode = jnp.where( update_L_mode_values, - core_profiles.T_e.value[ped_top_idx], # pyrefly: ignore[bad-index] + core_profiles.T_e.value[ped_top_idx], pedestal_transition_state.T_e_ped_L_mode, ) new_n_e_ped_L_mode = jnp.where( update_L_mode_values, - core_profiles.n_e.value[ped_top_idx], # pyrefly: ignore[bad-index] + core_profiles.n_e.value[ped_top_idx], pedestal_transition_state.n_e_ped_L_mode, ) @@ -432,8 +433,9 @@ def pre_step( n_e=input_state.core_sources.n_e | explicit_source_profiles.n_e, psi=input_state.core_sources.psi | explicit_source_profiles.psi, ) + assert pedestal_transition_state is not None pedestal_transition_state = _update_pedestal_transition_state( - pedestal_transition_state=pedestal_transition_state, # pyrefly: ignore[bad-argument-type] + pedestal_transition_state=pedestal_transition_state, runtime_params=runtime_params_t, geo=geo_t, core_profiles=input_state.core_profiles, @@ -445,14 +447,15 @@ def pre_step( # and freeze its output for the solver loop. calc_coeffs will skip # re-evaluation and use this stored output. if runtime_params_t.pedestal.explicit_pedestal: + assert pedestal_transition_state is not None pedestal_model_output = models.pedestal_model( runtime_params_t, geo_t, input_state.core_profiles, explicit_source_profiles, - pedestal_transition_state, # pyrefly: ignore[bad-argument-type] + pedestal_transition_state, ) - pedestal_transition_state = dataclasses.replace( # pyrefly: ignore[bad-specialization] + pedestal_transition_state = dataclasses.replace( pedestal_transition_state, pedestal_model_output=pedestal_model_output, ) diff --git a/torax/_src/output_tools/output.py b/torax/_src/output_tools/output.py index 13bc5b124..f6f40d75f 100644 --- a/torax/_src/output_tools/output.py +++ b/torax/_src/output_tools/output.py @@ -594,7 +594,7 @@ def _save_geometry( field_name = output_keys.Z_MAGNETIC_AXIS data_array = self._pack_into_data_array( field_name, - data, # pyrefly: ignore[bad-argument-type] + data, ) if data_array is not None: xr_dict[field_name] = data_array @@ -636,7 +636,7 @@ def _save_geometry( # _face variables with no corresponding non-face variable. if name.endswith("_face"): name = name.removesuffix("_face") - data_array = self._pack_into_data_array(name, property_data) # pyrefly: ignore[bad-argument-type] + data_array = self._pack_into_data_array(name, property_data) if data_array is not None: xr_dict[name] = data_array diff --git a/torax/_src/output_tools/post_processing.py b/torax/_src/output_tools/post_processing.py index 404f3692b..b71e0b55f 100644 --- a/torax/_src/output_tools/post_processing.py +++ b/torax/_src/output_tools/post_processing.py @@ -15,7 +15,7 @@ """Functions for adding post-processed outputs to the simulation state.""" import dataclasses -from typing import Callable +from typing import Callable, Self from absl import logging import jax @@ -36,7 +36,6 @@ from torax._src.physics import rotation from torax._src.physics import scaling_laws from torax._src.sources import source_profiles -import typing_extensions # pylint: disable=invalid-name @@ -330,7 +329,7 @@ class PostProcessedOutputs: # pylint: enable=invalid-name @classmethod - def zeros(cls, geo: geometry.Geometry) -> typing_extensions.Self: + def zeros(cls, geo: geometry.Geometry) -> Self: """Returns a PostProcessedOutputs with all zeros, used for initializing.""" return cls( pprime=jnp.zeros(geo.rho_face.shape), @@ -680,7 +679,7 @@ def make_post_processed_outputs( ) # Calculate normalized poloidal flux. psi_face = sim_state.core_profiles.psi.face_value() - psi_norm_face = (psi_face - psi_face[0]) / (psi_face[-1] - psi_face[0]) # pyrefly: ignore[bad-index] + psi_norm_face = (psi_face - psi_face[0]) / (psi_face[-1] - psi_face[0]) integrated_sources = _calculate_integrated_sources( sim_state.geometry, sim_state.core_sources, @@ -821,24 +820,24 @@ def cumulative_values(): # Calculate te and ti volume average [keV] te_volume_avg = math_utils.volume_average( - sim_state.core_profiles.T_e.value, sim_state.geometry # pyrefly: ignore[bad-argument-type] + sim_state.core_profiles.T_e.value, sim_state.geometry ) ti_volume_avg = math_utils.volume_average( - sim_state.core_profiles.T_i.value, sim_state.geometry # pyrefly: ignore[bad-argument-type] + sim_state.core_profiles.T_i.value, sim_state.geometry ) # Calculate n_e and n_i (main ion) volume and line averages in m^-3 n_e_volume_avg = math_utils.volume_average( - sim_state.core_profiles.n_e.value, sim_state.geometry # pyrefly: ignore[bad-argument-type] + sim_state.core_profiles.n_e.value, sim_state.geometry ) n_i_volume_avg = math_utils.volume_average( - sim_state.core_profiles.n_i.value, sim_state.geometry # pyrefly: ignore[bad-argument-type] + sim_state.core_profiles.n_i.value, sim_state.geometry ) n_e_line_avg = math_utils.line_average( - sim_state.core_profiles.n_e.value, sim_state.geometry # pyrefly: ignore[bad-argument-type] + sim_state.core_profiles.n_e.value, sim_state.geometry ) n_i_line_avg = math_utils.line_average( - sim_state.core_profiles.n_i.value, sim_state.geometry # pyrefly: ignore[bad-argument-type] + sim_state.core_profiles.n_i.value, sim_state.geometry ) fgw_n_e_volume_avg = formulas.calculate_greenwald_fraction( n_e_volume_avg, sim_state.core_profiles, sim_state.geometry @@ -905,7 +904,7 @@ def cumulative_values(): runtime_params.numerics.min_rho_norm, ) j_toroidal_external = psi_calculations.j_parallel_to_j_toroidal( - j_parallel_external, # pyrefly: ignore[bad-argument-type] + j_parallel_external, sim_state.geometry, runtime_params.numerics.min_rho_norm, ) @@ -1017,7 +1016,7 @@ def cumulative_values(): j_ecrh=j_toroidal_sources['j_ecrh'], j_generic_current=j_toroidal_sources['j_generic_current'], j_non_inductive=j_toroidal_bootstrap + j_toroidal_external, - j_parallel_external=j_parallel_external, # pyrefly: ignore[bad-argument-type] + j_parallel_external=j_parallel_external, j_parallel_non_inductive=j_parallel_bootstrap + j_parallel_external, I_external=I_external, I_non_inductive=I_non_inductive, @@ -1037,8 +1036,8 @@ def cumulative_values(): beta_pol_profile=beta_pol_profile.face_value(), beta_pol_prime=beta_pol_prime, impurity_species=impurity_radiation_outputs, - poloidal_velocity=rotation_output.poloidal_velocity.face_value(), # pyrefly: ignore[bad-argument-type] - radial_electric_field=rotation_output.Er.face_value(), # pyrefly: ignore[bad-argument-type] + poloidal_velocity=rotation_output.poloidal_velocity.face_value(), + radial_electric_field=rotation_output.Er.face_value(), first_step=jnp.array(False), ) diff --git a/torax/_src/output_tools/tests/post_processing_test.py b/torax/_src/output_tools/tests/post_processing_test.py index 9a42423e9..dfb423f86 100644 --- a/torax/_src/output_tools/tests/post_processing_test.py +++ b/torax/_src/output_tools/tests/post_processing_test.py @@ -58,23 +58,23 @@ def setUp(self): ), qei=source_profiles_lib.QeiInfo.zeros(self.geo), T_i={ # pyrefly: ignore[bad-argument-type] - 'fusion': ones, # pyrefly: ignore[bad-assignment] - 'generic_heat': 2 * ones, # pyrefly: ignore[bad-assignment] - 'icrh': 3 * ones, # pyrefly: ignore[bad-assignment] + 'fusion': ones, + 'generic_heat': 2 * ones, + 'icrh': 3 * ones, }, T_e={ # pyrefly: ignore[bad-argument-type] - 'bremsstrahlung': -ones, # pyrefly: ignore[bad-assignment] - 'cyclotron_radiation': -2 * ones, # pyrefly: ignore[bad-assignment] - 'impurity_radiation': -3 * ones, # pyrefly: ignore[bad-assignment] - 'ohmic': 5 * ones, # pyrefly: ignore[bad-assignment] - 'fusion': ones, # pyrefly: ignore[bad-assignment] - 'generic_heat': 3 * ones, # pyrefly: ignore[bad-assignment] - 'ecrh': 7 * ones, # pyrefly: ignore[bad-assignment] - 'icrh': 1.5 * ones, # pyrefly: ignore[bad-assignment] + 'bremsstrahlung': -ones, + 'cyclotron_radiation': -2 * ones, + 'impurity_radiation': -3 * ones, + 'ohmic': 5 * ones, + 'fusion': ones, + 'generic_heat': 3 * ones, + 'ecrh': 7 * ones, + 'icrh': 1.5 * ones, }, psi={ # pyrefly: ignore[bad-argument-type] - 'generic_current': 2 * ones, # pyrefly: ignore[bad-assignment] - 'ecrh': 2 * ones, # pyrefly: ignore[bad-assignment] + 'generic_current': 2 * ones, + 'ecrh': 2 * ones, }, n_e={}, ) diff --git a/torax/_src/pedestal_model/formation/power_scaling_formation_model.py b/torax/_src/pedestal_model/formation/power_scaling_formation_model.py index 5f0c2d52a..511b75146 100644 --- a/torax/_src/pedestal_model/formation/power_scaling_formation_model.py +++ b/torax/_src/pedestal_model/formation/power_scaling_formation_model.py @@ -43,11 +43,11 @@ class PowerScalingFormationRuntimeParams( def calculate_P_SOL_total( - internal_plasma_energy: state.PlasmaInternalEnergy, + internal_plasma_energy: state.PlasmaInternalEnergy | None, core_sources: source_profiles_lib.SourceProfiles, geo: geometry.Geometry, include_dW_dt: bool = True, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the total power out of the separatrix. Args: @@ -69,9 +69,9 @@ def calculate_P_SOL_total( for source in core_sources.T_i.values() ) P_heat_total = P_heat_e + P_heat_i - if not include_dW_dt: - return P_heat_total # pyrefly: ignore[bad-return] - return P_heat_total - internal_plasma_energy.dW_thermal_dt_smoothed # pyrefly: ignore[bad-return] + if not include_dW_dt or internal_plasma_energy is None: + return P_heat_total + return P_heat_total - internal_plasma_energy.dW_thermal_dt_smoothed @dataclasses.dataclass(frozen=True, eq=False) @@ -103,7 +103,7 @@ def __call__( ) P_SOL_total = calculate_P_SOL_total( - core_profiles.internal_plasma_energy, # pyrefly: ignore[bad-argument-type] + core_profiles.internal_plasma_energy, core_sources, geo, include_dW_dt=runtime_params.pedestal.include_dW_dt_in_P_SOL, diff --git a/torax/_src/pedestal_model/pedestal_model_output.py b/torax/_src/pedestal_model/pedestal_model_output.py index 67cec1a13..2c029e402 100644 --- a/torax/_src/pedestal_model/pedestal_model_output.py +++ b/torax/_src/pedestal_model/pedestal_model_output.py @@ -194,8 +194,8 @@ def _tanh_internal_boundary_conditions( """ # Get ψ_N at each cell grid point. psi_face = core_profiles.psi.face_value() - psi_norm_cell = (core_profiles.psi.value - psi_face[0]) / ( # pyrefly: ignore[bad-index] - psi_face[-1] - psi_face[0] # pyrefly: ignore[bad-index] + psi_norm_cell = (core_profiles.psi.value - psi_face[0]) / ( + psi_face[-1] - psi_face[0] ) # Derive Δ from rho_norm_ped_top via ψ_N mapping. diff --git a/torax/_src/pedestal_model/pydantic_model.py b/torax/_src/pedestal_model/pydantic_model.py index e7cbb3d3a..8f11762c1 100644 --- a/torax/_src/pedestal_model/pydantic_model.py +++ b/torax/_src/pedestal_model/pydantic_model.py @@ -16,7 +16,7 @@ import abc import copy -from typing import Annotated, Any, Literal, TypeAlias +from typing import Annotated, Any, Literal, Self, TypeAlias import chex import pydantic from torax._src import array_typing @@ -29,7 +29,6 @@ from torax._src.pedestal_model.saturation import profile_value_saturation_model from torax._src.physics import scaling_laws from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # pylint: disable=invalid-name @@ -326,7 +325,7 @@ def _defaults(cls, data: dict[str, Any]) -> dict[str, Any]: return configurable_data @pydantic.model_validator(mode="after") - def _check_source_mode(self) -> typing_extensions.Self: + def _check_source_mode(self) -> Self: if ( self.use_formation_model_with_internal_boundary_condition and self.mode != runtime_params.Mode.INTERNAL_BOUNDARY_CONDITION diff --git a/torax/_src/pedestal_model/saturation/profile_value_saturation_model.py b/torax/_src/pedestal_model/saturation/profile_value_saturation_model.py index 3abe31d99..da8e9888b 100644 --- a/torax/_src/pedestal_model/saturation/profile_value_saturation_model.py +++ b/torax/_src/pedestal_model/saturation/profile_value_saturation_model.py @@ -38,7 +38,7 @@ def __call__( geo: geometry.Geometry, core_profiles: state.CoreProfiles, pedestal_output: pedestal_model_output.PedestalModelOutput, - ) -> array_typing.FloatScalar: + ) -> pedestal_model_output.TransportMultipliers: """Calculates transport increase multipliers.""" # Get the current profile values at the pedestal top. # Interpolating to get the values at exactly rho_norm_ped_top is difficult, @@ -48,10 +48,10 @@ def __call__( rho_norm_face_ped_top_idx = jnp.argmin( jnp.abs(geo.rho_face_norm - pedestal_output.rho_norm_ped_top) ) - current_T_e_ped_top = core_profiles.T_e.face_value()[ # pyrefly: ignore[bad-index] + current_T_e_ped_top = core_profiles.T_e.face_value()[ rho_norm_face_ped_top_idx ] - current_T_i_ped_top = core_profiles.T_i.face_value()[ # pyrefly: ignore[bad-index] + current_T_i_ped_top = core_profiles.T_i.face_value()[ rho_norm_face_ped_top_idx ] @@ -63,7 +63,7 @@ def __call__( current_T_i_ped_top, pedestal_output.T_i_ped, runtime_params.pedestal ) - return pedestal_model_output.TransportMultipliers( # pyrefly: ignore[bad-return] + return pedestal_model_output.TransportMultipliers( chi_e_multiplier=chi_e_multiplier, chi_i_multiplier=chi_i_multiplier, # TODO(b/487920703): set the density transport coefficients based on diff --git a/torax/_src/pedestal_model/saturation/tests/profile_value_saturation_model_test.py b/torax/_src/pedestal_model/saturation/tests/profile_value_saturation_model_test.py index 742fc6a5b..40c489f4e 100644 --- a/torax/_src/pedestal_model/saturation/tests/profile_value_saturation_model_test.py +++ b/torax/_src/pedestal_model/saturation/tests/profile_value_saturation_model_test.py @@ -66,7 +66,7 @@ def test_saturation_multiplier( # For this test, we put the pedestal top at the last grid point. ped_top_idx = -1 - current_T_e_ped = self.core_profiles.T_e.face_value()[ped_top_idx] # pyrefly: ignore[bad-index] + current_T_e_ped = self.core_profiles.T_e.face_value()[ped_top_idx] # Construct a pedestal output that is asking for a pedestal with # target temperature. diff --git a/torax/_src/pedestal_model/set_pped_tpedratio_nped.py b/torax/_src/pedestal_model/set_pped_tpedratio_nped.py index 9ba9ab946..e5eff72c5 100644 --- a/torax/_src/pedestal_model/set_pped_tpedratio_nped.py +++ b/torax/_src/pedestal_model/set_pped_tpedratio_nped.py @@ -14,6 +14,7 @@ """Pedestal model that specifies pressure, temperature ratio, and density.""" import dataclasses +from typing import override import jax from jax import numpy as jnp @@ -27,7 +28,6 @@ from torax._src.pedestal_model import pedestal_transition_state as pedestal_transition_state_lib from torax._src.pedestal_model import runtime_params as pedestal_runtime_params_lib from torax._src.physics import formulas -from typing_extensions import override # pylint: disable=invalid-name diff --git a/torax/_src/pedestal_model/set_tped_nped.py b/torax/_src/pedestal_model/set_tped_nped.py index 8c460aa33..7fa4592ef 100644 --- a/torax/_src/pedestal_model/set_tped_nped.py +++ b/torax/_src/pedestal_model/set_tped_nped.py @@ -14,6 +14,7 @@ """A basic version of the pedestal model that uses direct specification.""" import dataclasses +from typing import override import jax from jax import numpy as jnp @@ -25,7 +26,6 @@ from torax._src.pedestal_model import pedestal_model_output from torax._src.pedestal_model import pedestal_transition_state as pedestal_transition_state_lib from torax._src.pedestal_model import runtime_params as pedestal_runtime_params_lib -from typing_extensions import override # pylint: disable=invalid-name diff --git a/torax/_src/pedestal_model/tests/pedestal_model_output_test.py b/torax/_src/pedestal_model/tests/pedestal_model_output_test.py index 78961528e..4d41ebe8a 100644 --- a/torax/_src/pedestal_model/tests/pedestal_model_output_test.py +++ b/torax/_src/pedestal_model/tests/pedestal_model_output_test.py @@ -281,7 +281,7 @@ def _make_standard_cell_var(cell_vals, right_face_val): # Compute psi_norm at cell centers using face_value() to match impl. psi_fv = core_profiles.psi.face_value() - psi_norm_cell = (psi_cell - psi_fv[0]) / (psi_fv[-1] - psi_fv[0]) # pyrefly: ignore[bad-index] + psi_norm_cell = (psi_cell - psi_fv[0]) / (psi_fv[-1] - psi_fv[0]) # Derive delta from nearest cell to rho_ped_top (matching impl). ped_top_idx = jnp.argmin(jnp.abs(geo.rho_norm - rho_ped_top)) diff --git a/torax/_src/physics/collisions.py b/torax/_src/physics/collisions.py index 6b0e78710..d7e4681b0 100644 --- a/torax/_src/physics/collisions.py +++ b/torax/_src/physics/collisions.py @@ -34,6 +34,7 @@ Z=1 plasma. """ +import chex import jax from jax import numpy as jnp from torax._src import array_typing @@ -47,7 +48,7 @@ def coll_exchange( core_profiles: state.CoreProfiles, Qei_multiplier: float, -) -> jax.Array: +) -> array_typing.FloatVectorCell: """Computes collisional ion-electron heat exchange coefficient (equipartion). Args: @@ -59,13 +60,13 @@ def coll_exchange( """ # Calculate Coulomb logarithm log_lambda_ei = calculate_log_lambda_ei( - core_profiles.T_e.value, core_profiles.n_e.value # pyrefly: ignore[bad-argument-type] + core_profiles.T_e.value, core_profiles.n_e.value ) # ion-electron collisionality for Z_eff=1. Ion charge and multiple ion effects # are included in the Qei_coef calculation below. log_tau_e_Z1 = _calculate_log_tau_e_Z1( - core_profiles.T_e.value, # pyrefly: ignore[bad-argument-type] - core_profiles.n_e.value, # pyrefly: ignore[bad-argument-type] + core_profiles.T_e.value, + core_profiles.n_e.value, log_lambda_ei, ) # pylint: disable=invalid-name @@ -87,7 +88,7 @@ def calc_nu_star( geo: geometry.Geometry, core_profiles: state.CoreProfiles, collisionality_multiplier: float, -) -> jax.Array: +) -> array_typing.FloatVectorFace: """Calculates nustar. Electron-ion collision frequency normalized by bounce frequency. @@ -104,14 +105,14 @@ def calc_nu_star( # Calculate Coulomb logarithm log_lambda_ei_face = calculate_log_lambda_ei( - core_profiles.T_e.face_value(), # pyrefly: ignore[bad-argument-type] - core_profiles.n_e.face_value(), # pyrefly: ignore[bad-argument-type] + core_profiles.T_e.face_value(), + core_profiles.n_e.face_value(), ) # ion_electron collisionality log_tau_e_Z1 = _calculate_log_tau_e_Z1( - core_profiles.T_e.face_value(), # pyrefly: ignore[bad-argument-type] - core_profiles.n_e.face_value(), # pyrefly: ignore[bad-argument-type] + core_profiles.T_e.face_value(), + core_profiles.n_e.face_value(), log_lambda_ei_face, ) @@ -123,7 +124,7 @@ def calc_nu_star( ) # calculate bounce time - tau_bounce = ( + tau_bounce = jnp.asarray( core_profiles.q_face * geo.R_major_profile_face / ( @@ -136,7 +137,7 @@ def calc_nu_star( ) ) # due to pathological on-axis epsilon=0 term - tau_bounce = tau_bounce.at[0].set(tau_bounce[1]) # pyrefly: ignore[missing-attribute] + tau_bounce = tau_bounce.at[0].set(tau_bounce[1]) # calculate normalized collisionality nustar = nu_e * tau_bounce @@ -182,9 +183,9 @@ def fast_ion_fractional_heating_formula( def calculate_log_lambda_ee( - T_e: jax.Array, - n_e: jax.Array, -) -> jax.Array: + T_e: chex.Numeric, + n_e: chex.Numeric, +) -> array_typing.Array: """Calculates Coulomb logarithm for electron-electron collisions. Note: the difference with calculate_log_lambda_ei is minimal. @@ -204,9 +205,9 @@ def calculate_log_lambda_ee( def calculate_log_lambda_ei( - T_e: jax.Array, - n_e: jax.Array, -) -> jax.Array: + T_e: chex.Numeric, + n_e: chex.Numeric, +) -> array_typing.Array: """Calculates Coulomb logarithm for electron-ion collisions. See Wesson 3rd edition p727. @@ -224,10 +225,10 @@ def calculate_log_lambda_ei( def calculate_log_lambda_ii( - T_i: jax.Array, - n_i: jax.Array, - Z_i: jax.Array, -) -> jax.Array: + T_i: chex.Numeric, + n_i: chex.Numeric, + Z_i: chex.Numeric, +) -> array_typing.Array: """Calculates Coulomb logarithm for ion-ion collisions. Formula 18e in Sauter PoP 1999. See also NRL formulary 2013, page 34. @@ -246,12 +247,12 @@ def calculate_log_lambda_ii( def calculate_tau_ii( - A_i: jax.Array, - Z_i: jax.Array, - T_i: jax.Array, - n_i: jax.Array, - ln_Lambda_ii: jax.Array, -) -> jax.Array: + A_i: chex.Numeric, + Z_i: chex.Numeric, + T_i: chex.Numeric, + n_i: chex.Numeric, + ln_Lambda_ii: chex.Numeric, +) -> array_typing.Array: """Calculates ion-ion (self) collision time for a single ion species. See Wesson 3rd edition p730. @@ -283,9 +284,9 @@ def calculate_tau_ii( # TODO(b/377225415): generalize to arbitrary number of ions. def _calculate_weighted_Z_eff( core_profiles: state.CoreProfiles, -) -> jax.Array: +) -> array_typing.FloatVectorCell: """Calculates ion mass weighted Z_eff. Used for collisional heat exchange.""" - return ( # pyrefly: ignore[bad-return] + return ( core_profiles.n_i.value * core_profiles.Z_i**2 / core_profiles.A_i + core_profiles.n_impurity.value * core_profiles.Z_impurity**2 @@ -294,10 +295,10 @@ def _calculate_weighted_Z_eff( def _calculate_log_tau_e_Z1( - T_e: jax.Array, - n_e: jax.Array, - log_lambda_ei: jax.Array, -) -> jax.Array: + T_e: chex.Numeric, + n_e: chex.Numeric, + log_lambda_ei: chex.Numeric, +) -> array_typing.Array: """Calculates log of electron-ion collision time for Z=1 plasma. See Wesson 3rd edition p729. Extension to multiple ions is context dependent diff --git a/torax/_src/physics/fast_ion_utils.py b/torax/_src/physics/fast_ion_utils.py index 008a91fe8..cf37810a6 100644 --- a/torax/_src/physics/fast_ion_utils.py +++ b/torax/_src/physics/fast_ion_utils.py @@ -14,8 +14,10 @@ """Fast ion utility functions.""" +import chex import jax from jax import numpy as jnp +from torax._src import array_typing from torax._src import constants from torax._src import math_utils from torax._src.physics import collisions @@ -25,14 +27,14 @@ def _nu_epsilon( - m_a_amu: float, - Z_a: float, - T_a_keV: jax.Array, - m_b_amu: float, - Z_b: float, - n_b_m3: jax.Array, - T_b_keV: jax.Array, - ln_lambda: jax.Array, + m_a_amu: chex.Numeric, + Z_a: chex.Numeric, + T_a_keV: array_typing.Array, + m_b_amu: chex.Numeric, + Z_b: chex.Numeric, + n_b_m3: array_typing.Array, + T_b_keV: array_typing.Array, + ln_lambda: array_typing.Array, ) -> jax.Array: """NRL Formulary energy exchange rate nu_epsilon [Hz]. @@ -73,13 +75,13 @@ def _nu_epsilon( def _compute_T_tail( - P_density_W: jax.Array, - T_e: jax.Array, - n_e: jax.Array, - n_total: jax.Array, - charge_number: float, - mass_number: float, -) -> jax.Array: + P_density_W: array_typing.Array, + T_e: array_typing.Array, + n_e: array_typing.Array, + n_total: array_typing.Array, + charge_number: chex.Numeric, + mass_number: chex.Numeric, +) -> array_typing.Array: """Computes the effective tail temperature via the Stix xi parameter. Uses the Spitzer slowing-down time on electrons (tau_s) and the Stix @@ -124,20 +126,20 @@ def _compute_T_tail( def bimaxwellian_split( - power_deposition: jax.Array, - T_e: jax.Array, - n_e: jax.Array, - T_i: jax.Array, - n_i: jax.Array, - minority_concentration: jax.Array | float, - P_total_W: float, - charge_number: float, - mass_number: float, - bulk_ion_mass: float, - Z_i: float, - n_impurity: jax.Array, - Z_impurity: float, - A_impurity: float, + power_deposition: array_typing.Array, + T_e: array_typing.Array, + n_e: array_typing.Array, + T_i: array_typing.Array, + n_i: array_typing.Array, + minority_concentration: chex.Numeric, + P_total_W: chex.Numeric, + charge_number: chex.Numeric, + mass_number: chex.Numeric, + bulk_ion_mass: chex.Numeric, + Z_i: chex.Numeric, + n_impurity: array_typing.Array, + Z_impurity: chex.Numeric, + A_impurity: chex.Numeric, ) -> tuple[jax.Array, jax.Array]: """Returns (n_tail, T_tail) using the Power Balance Closure. @@ -194,7 +196,7 @@ def bimaxwellian_split( mass_number, charge_number, T_tail, - me_amu, # pyrefly: ignore[bad-argument-type] + me_amu, 1.0, n_e, T_e, diff --git a/torax/_src/physics/formulas.py b/torax/_src/physics/formulas.py index a4b0c8a63..76ad3efc4 100644 --- a/torax/_src/physics/formulas.py +++ b/torax/_src/physics/formulas.py @@ -169,9 +169,9 @@ def calculate_stored_thermal_energy( wth_ion: Ion thermal stored energy [J] wth_tot: Total thermal stored energy [J] """ - wth_el = math_utils.volume_integration(1.5 * p_el.value, geo) # pyrefly: ignore[bad-argument-type] - wth_ion = math_utils.volume_integration(1.5 * p_ion.value, geo) # pyrefly: ignore[bad-argument-type] - wth_tot = math_utils.volume_integration(1.5 * p_tot.value, geo) # pyrefly: ignore[bad-argument-type] + wth_el = math_utils.volume_integration(1.5 * p_el.value, geo) + wth_ion = math_utils.volume_integration(1.5 * p_ion.value, geo) + wth_tot = math_utils.volume_integration(1.5 * p_tot.value, geo) return wth_el, wth_ion, wth_tot @@ -205,7 +205,9 @@ def calculate_greenwald_fraction( def calculate_betas( core_profiles: state.CoreProfiles, geo: geometry.Geometry, -) -> array_typing.FloatScalar: +) -> tuple[ + array_typing.FloatScalar, array_typing.FloatScalar, array_typing.FloatScalar +]: """Calculates the beta_tor, beta_pol, and beta_N plasma beta quantities. beta_tor is defined as the ratio of volume-averaged plasma pressure to @@ -237,13 +239,13 @@ def calculate_betas( Tuple of beta_tor, beta_pol, and beta_N """ p_total_volume_avg = math_utils.volume_average( - core_profiles.pressure_total.value, geo # pyrefly: ignore[bad-argument-type] + core_profiles.pressure_total.value, geo ) magnetic_pressure_on_axis = geo.B_0**2 / (2 * constants.CONSTANTS.mu_0) # Add a division guard though B0 should typically be non-zero. beta_tor = math_utils.safe_divide( - num=p_total_volume_avg, denom=magnetic_pressure_on_axis, eps=1e-7 # pyrefly: ignore[bad-argument-type] + num=p_total_volume_avg, denom=magnetic_pressure_on_axis, eps=1e-7 ) beta_pol = ( @@ -268,7 +270,7 @@ def calculate_betas( ) ) - return beta_tor, beta_pol, beta_N # pyrefly: ignore[bad-return] + return beta_tor, beta_pol, beta_N def calculate_beta_pol_profile( diff --git a/torax/_src/physics/psi_calculations.py b/torax/_src/physics/psi_calculations.py index b55d3c700..0f77fd4be 100644 --- a/torax/_src/physics/psi_calculations.py +++ b/torax/_src/physics/psi_calculations.py @@ -112,12 +112,12 @@ def calc_q_face( """Calculates the q-profile on the face grid given poloidal flux (psi).""" # iota is standard terminology for 1/q inv_iota = jnp.abs( - (2 * geo.Phi_b * geo.rho_face_norm[1:]) / psi.face_grad()[1:] # pyrefly: ignore[bad-index] + (2 * geo.Phi_b * geo.rho_face_norm[1:]) / psi.face_grad()[1:] ) # Use L'Hôpital's rule to calculate iota on-axis, with psi_face_grad()[0]=0. inv_iota0 = jnp.expand_dims( - jnp.abs((2 * geo.Phi_b * geo.drho_norm[0]) / psi.face_grad()[1]), 0 # pyrefly: ignore[bad-index] + jnp.abs((2 * geo.Phi_b * geo.drho_norm[0]) / psi.face_grad()[1]), 0 ) q_face = jnp.concatenate([inv_iota0, inv_iota]) @@ -191,33 +191,34 @@ def calc_j_total( def calc_s_face( geo: geometry.Geometry, psi: cell_variable.CellVariable -) -> jax.Array: +) -> array_typing.FloatVectorFace: """Calculates magnetic shear on the face grid from poloidal flux (psi).""" # iota (1/q) should have a /2*Phib but we drop it since will cancel out in # the s calculation. - iota_scaled = jnp.abs((psi.face_grad()[1:] / geo.rho_face_norm[1:])) # pyrefly: ignore[bad-index] + iota_scaled = jnp.abs((psi.face_grad()[1:] / geo.rho_face_norm[1:])) # on-axis iota_scaled from L'Hôpital's rule = dpsi_face_grad / drho_norm # Using expand_dims to make it compatible with jnp.concatenate iota_scaled0 = jnp.expand_dims( - jnp.abs(psi.face_grad()[1] / geo.drho_norm[0]), axis=0 # pyrefly: ignore[bad-index] + jnp.abs(psi.face_grad()[1] / geo.drho_norm[0]), axis=0 ) iota_scaled = jnp.concatenate([iota_scaled0, iota_scaled]) + grad_iota = jnp.asarray(jnp.gradient(iota_scaled, geo.rho_face_norm)) s_face = ( - -geo.rho_face_norm # pyrefly: ignore[unsupported-operation] - * jnp.gradient(iota_scaled, geo.rho_face_norm) + -1 * geo.rho_face_norm + * grad_iota / iota_scaled ) - return s_face # pyrefly: ignore[bad-return] + return s_face def calc_s_rmid( geo: geometry.Geometry, psi: cell_variable.CellVariable -) -> jax.Array: +) -> array_typing.FloatVectorFace: """Calculates magnetic shear (s) from poloidal flux (psi). Version taking the derivative of iota with respect to the midplane r, @@ -233,21 +234,22 @@ def calc_s_rmid( # iota (1/q) should have a /2*Phib but we drop it since will cancel out in # the s calculation. - iota_scaled = jnp.abs((psi.face_grad()[1:] / geo.rho_face_norm[1:])) # pyrefly: ignore[bad-index] + iota_scaled = jnp.abs((psi.face_grad()[1:] / geo.rho_face_norm[1:])) # on-axis iota_scaled from L'Hôpital's rule = dpsi_face_grad / drho_norm # Using expand_dims to make it compatible with jnp.concatenate iota_scaled0 = jnp.expand_dims( - jnp.abs(psi.face_grad()[1] / geo.drho_norm[0]), axis=0 # pyrefly: ignore[bad-index] + jnp.abs(psi.face_grad()[1] / geo.drho_norm[0]), axis=0 ) iota_scaled = jnp.concatenate([iota_scaled0, iota_scaled]) rmid_face = (geo.R_out_face - geo.R_in_face) * 0.5 - s_face = -rmid_face * jnp.gradient(iota_scaled, rmid_face) / iota_scaled # pyrefly: ignore[unsupported-operation] + grad_iota = jnp.asarray(jnp.gradient(iota_scaled, rmid_face)) + s_face = -1 * rmid_face * grad_iota / iota_scaled - return s_face # pyrefly: ignore[bad-return] + return s_face def calc_bpol_squared( @@ -268,7 +270,7 @@ def calc_bpol_squared( bpol2_face: Square of poloidal magnetic field, on the face grid. """ bpol2_bulk = ( - (psi.face_grad()[1:] / (2 * jnp.pi)) ** 2 # pyrefly: ignore[bad-index] + (psi.face_grad()[1:] / (2 * jnp.pi)) ** 2 * geo.g2_face[1:] / geo.vpr_face[1:] ** 2 ) @@ -338,10 +340,10 @@ def calc_q95( def calculate_psi_grad_constraint_from_Ip( Ip: array_typing.FloatScalar, geo: geometry.Geometry, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the gradient constraint on the poloidal flux (psi) from Ip.""" return ( - Ip # pyrefly: ignore[bad-return] + Ip * (16 * jnp.pi**3 * constants.CONSTANTS.mu_0 * geo.Phi_b) / (geo.g2g3_over_rhon_face[-1] * geo.F_face[-1]) ) @@ -353,12 +355,12 @@ def calculate_psi_value_constraint_from_v_loop( v_loop_lcfs_t: array_typing.FloatScalar, v_loop_lcfs_t_plus_dt: array_typing.FloatScalar, psi_lcfs_t: array_typing.FloatScalar, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the value constraint on the poloidal flux for the next time step from loop voltage.""" theta_weighted_v_loop_lcfs = ( 1 - theta ) * v_loop_lcfs_t + theta * v_loop_lcfs_t_plus_dt - return psi_lcfs_t + theta_weighted_v_loop_lcfs * dt # pyrefly: ignore[bad-return] + return psi_lcfs_t + theta_weighted_v_loop_lcfs * dt # TODO(b/406173731): Find robust solution for underdetermination and solve this @@ -367,7 +369,7 @@ def calculate_v_loop_lcfs_from_psi( psi_t: cell_variable.CellVariable, psi_t_plus_dt: cell_variable.CellVariable, dt: array_typing.FloatScalar, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the v_loop_lcfs for the next timestep. For the Ip boundary condition case, the v_loop_lcfs formula is in principle @@ -389,8 +391,8 @@ def calculate_v_loop_lcfs_from_psi( Returns: The updated v_loop_lcfs for the next timestep. """ - psi_lcfs_t = psi_t.face_value()[-1] # pyrefly: ignore[bad-index] - psi_lcfs_t_plus_dt = psi_t_plus_dt.face_value()[-1] # pyrefly: ignore[bad-index] + psi_lcfs_t = psi_t.face_value()[-1] + psi_lcfs_t_plus_dt = psi_t_plus_dt.face_value()[-1] v_loop_lcfs_t_plus_dt = (psi_lcfs_t_plus_dt - psi_lcfs_t) / dt return v_loop_lcfs_t_plus_dt @@ -399,10 +401,10 @@ def calculate_psidot_from_psi_sources( *, psi_sources: array_typing.FloatVector, sigma: array_typing.FloatVector, - resistivity_multiplier: float, + resistivity_multiplier: array_typing.FloatScalar, psi: cell_variable.CellVariable, geo: geometry.Geometry, -) -> jax.Array: +) -> array_typing.FloatVectorCell: """Calculates psidot (loop voltage) from the sum of the psi sources.""" # Calculate transient term @@ -439,13 +441,13 @@ def calculate_psidot_from_psi_sources( d_face_psi, psi ) conv_mat, conv_vec = convection_terms.make_convection_terms( - v_face_psi, d_face_psi, psi # pyrefly: ignore[bad-argument-type] + v_face_psi, d_face_psi, psi ) c_mat = diffusion_mat + conv_mat c = diffusion_vec + conv_vec + psi_sources - return (c_mat.matvec(psi.value) + c) / toc_psi # pyrefly: ignore[bad-argument-type, bad-return] + return (c_mat.matvec(psi.value) + c) / toc_psi def j_toroidal_to_j_parallel( diff --git a/torax/_src/physics/rotation.py b/torax/_src/physics/rotation.py index 1295cfb4c..526f8623c 100644 --- a/torax/_src/physics/rotation.py +++ b/torax/_src/physics/rotation.py @@ -103,7 +103,7 @@ def _calculate_radial_electric_field( ) Er_toroidal_face = ( - -toroidal_angular_velocity.face_value() # pyrefly: ignore[unsupported-operation] + -toroidal_angular_velocity.face_value() * geo.R_major_profile_face * B_pol_face ) @@ -196,7 +196,7 @@ def calculate_rotation( ) ) - v_ExB = _calculate_v_ExB(Er.face_value(), B_total_face) # pyrefly: ignore[bad-argument-type] + v_ExB = _calculate_v_ExB(Er.face_value(), B_total_face) v_ExB_poloidal_and_pressure = _calculate_v_ExB( Er_poloidal_and_pressure_face, B_total_face ) diff --git a/torax/_src/physics/scaling_laws.py b/torax/_src/physics/scaling_laws.py index a440a32d0..4ca1b2a7b 100644 --- a/torax/_src/physics/scaling_laws.py +++ b/torax/_src/physics/scaling_laws.py @@ -25,6 +25,7 @@ import enum import jax from jax import numpy as jnp +from torax._src import array_typing from torax._src import math_utils from torax._src import state from torax._src.geometry import geometry @@ -72,20 +73,20 @@ class DivertorConfiguration(enum.StrEnum): class PLHAuxiliaryData: """Auxiliary data for P_LH calculations.""" - P_LH_high_density: jax.Array - P_LH_low_density: jax.Array - P_LH_min: jax.Array - line_average_n_e_at_P_LH_min: jax.Array + P_LH_high_density: array_typing.FloatScalar + P_LH_low_density: array_typing.FloatScalar + P_LH_min: array_typing.FloatScalar + line_average_n_e_at_P_LH_min: array_typing.FloatScalar def _calculate_P_LH_high_density( geo: geometry.Geometry, core_profiles: state.CoreProfiles, - line_average_n_e: jax.Array, + line_average_n_e: array_typing.FloatScalar, scaling_law: PLHScalingLaw, divertor_factor: float = 1.0, custom_prefactor: float = 1.0, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the H-mode transition power for the high density branch. See Eq. 3 in E. Delabie et al 2026 Nucl. Fusion 66 036016 for the general @@ -119,17 +120,17 @@ def _calculate_P_LH_high_density( * S ** params['S_exponent'] ) - return P_LH_MW * 1e6 # pyrefly: ignore[bad-return] + return P_LH_MW * 1e6 def _calculate_line_average_n_e_at_P_LH_min( geo: geometry.Geometry, core_profiles: state.CoreProfiles, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the density at P_LH_min from equation 3 in Ryter 2014.""" Ip_total = core_profiles.Ip_profile_face[..., -1] return ( - 0.7 # pyrefly: ignore[bad-return] + 0.7 * (Ip_total / 1e6) ** 0.34 * geo.a_minor**-0.95 * geo.B_0**0.62 @@ -144,7 +145,7 @@ def calculate_P_LH( scaling_law: PLHScalingLaw, divertor_configuration: DivertorConfiguration = DivertorConfiguration.HT, prefactor: float = 1.0, -) -> tuple[jax.Array, PLHAuxiliaryData]: +) -> tuple[array_typing.FloatScalar, PLHAuxiliaryData]: """Calculates the H-mode transition power from a given scaling law. Args: @@ -170,7 +171,7 @@ def calculate_P_LH( else: divertor_factor = 1.0 - line_average_n_e = math_utils.line_average(core_profiles.n_e.value, geo) # pyrefly: ignore[bad-argument-type] + line_average_n_e = math_utils.line_average(core_profiles.n_e.value, geo) line_average_n_e_at_P_LH_min = _calculate_line_average_n_e_at_P_LH_min( geo, core_profiles ) @@ -179,7 +180,7 @@ def calculate_P_LH( P_LH_high_density = _calculate_P_LH_high_density( geo=geo, core_profiles=core_profiles, - line_average_n_e=line_average_n_e, # pyrefly: ignore[bad-argument-type] + line_average_n_e=line_average_n_e, scaling_law=scaling_law, divertor_factor=divertor_factor, custom_prefactor=prefactor, @@ -216,9 +217,9 @@ def calculate_P_LH( def calculate_scaling_law_confinement_time( geo: geometry.Geometry, core_profiles: state.CoreProfiles, - P_loss: jax.Array, + P_loss: array_typing.FloatScalar, scaling_law: str, -) -> jax.Array: +) -> array_typing.FloatScalar: """Calculates the thermal energy confinement time for a given scaling law. Args: @@ -306,7 +307,7 @@ def calculate_scaling_law_confinement_time( scaled_Ploss = P_loss / 1e6 # convert to MW B = geo.B_0 line_avg_n_e = ( # convert to 10^19 m^-3 - math_utils.line_average(core_profiles.n_e.value, geo) / 1e19 # pyrefly: ignore[bad-argument-type] + math_utils.line_average(core_profiles.n_e.value, geo) / 1e19 ) R = geo.R_major inverse_aspect_ratio = geo.a_minor / geo.R_major diff --git a/torax/_src/physics/tests/formulas_test.py b/torax/_src/physics/tests/formulas_test.py index c6a73275d..36cae1f3c 100644 --- a/torax/_src/physics/tests/formulas_test.py +++ b/torax/_src/physics/tests/formulas_test.py @@ -78,9 +78,9 @@ def test_calculate_stored_thermal_energy(self): volume = math_utils.volume_integration(np.array([1.0]), self.geo) - np.testing.assert_allclose(wth_el, 1.5 * p_el.value[0] * volume) # pyrefly: ignore[bad-index] - np.testing.assert_allclose(wth_ion, 1.5 * p_ion.value[0] * volume) # pyrefly: ignore[bad-index] - np.testing.assert_allclose(wth_tot, 1.5 * p_tot.value[0] * volume) # pyrefly: ignore[bad-index] + np.testing.assert_allclose(wth_el, 1.5 * p_el.value[0] * volume) + np.testing.assert_allclose(wth_ion, 1.5 * p_ion.value[0] * volume) + np.testing.assert_allclose(wth_tot, 1.5 * p_tot.value[0] * volume) def test_calculate_greenwald_fraction(self): """Test that Greenwald fraction is calculated correctly.""" diff --git a/torax/_src/physics/tests/psi_calculations_test.py b/torax/_src/physics/tests/psi_calculations_test.py index e69f57e60..4a8443bdc 100644 --- a/torax/_src/physics/tests/psi_calculations_test.py +++ b/torax/_src/physics/tests/psi_calculations_test.py @@ -420,12 +420,12 @@ def test_update_v_loop_lcfs_from_psi(self): psi_t = cell_variable.CellVariable( value=np.ones_like(mesh.cell_centers) * 0.5, face_centers=mesh.face_centers, - right_face_grad_constraint=0.0, # pyrefly: ignore[bad-argument-type] + right_face_grad_constraint=0.0, ) psi_t_plus_dt = cell_variable.CellVariable( value=np.ones_like(mesh.cell_centers) * psi_lcfs_t_plus_dt, face_centers=mesh.face_centers, - right_face_grad_constraint=0.0, # pyrefly: ignore[bad-argument-type] + right_face_grad_constraint=0.0, ) v_loop_lcfs_t_plus_dt = psi_calculations.calculate_v_loop_lcfs_from_psi( diff --git a/torax/_src/sources/bremsstrahlung_heat_sink.py b/torax/_src/sources/bremsstrahlung_heat_sink.py index 62c769ede..59e33c748 100644 --- a/torax/_src/sources/bremsstrahlung_heat_sink.py +++ b/torax/_src/sources/bremsstrahlung_heat_sink.py @@ -21,6 +21,7 @@ import jax from jax import numpy as jnp import jaxtyping as jt +from torax._src import array_typing from torax._src import math_utils from torax._src import state from torax._src.config import runtime_params as runtime_params_lib @@ -50,7 +51,7 @@ def calc_bremsstrahlung( geo: geometry.Geometry, use_relativistic_correction: bool = False, exclude_impurity_bremsstrahlung: bool = False, -) -> tuple[jt.Float[jax.Array, ''], jt.Float[jax.Array, '']]: +) -> tuple[array_typing.FloatScalar, array_typing.FloatVectorCell]: """Calculate the Bremsstrahlung radiation power profile. Uses the model from Wesson, John, and David J. Campbell. Tokamaks. Vol. 149. @@ -85,18 +86,18 @@ def calc_bremsstrahlung( core_profiles.Z_eff_face, ) - P_brem_profile_face: jax.Array = ( + P_brem_profile_face: array_typing.FloatVectorFace = ( 5.35e-3 * Z_eff_face * n_e20**2 * jnp.sqrt(T_e_kev) ) # MW/m^3 - def calc_relativistic_correction() -> jax.Array: + def calc_relativistic_correction() -> array_typing.FloatVectorFace: # Apply the Stott relativistic correction. Tm = 511.0 # m_e * c**2 in keV correction = (1.0 + 2.0 * T_e_kev / Tm) * ( 1.0 + (2.0 / Z_eff_face) * (1.0 - 1.0 / (1.0 + T_e_kev / Tm)) ) - return correction # pyrefly: ignore[bad-return] + return correction # In MW/m^3 P_brem_profile_face = jnp.where( @@ -110,7 +111,7 @@ def calc_relativistic_correction() -> jax.Array: # In MW P_brem_total = math_utils.volume_integration(P_brem_profile_cell, geo) - return P_brem_total, P_brem_profile_cell # pyrefly: ignore[bad-return] + return P_brem_total, P_brem_profile_cell def bremsstrahlung_model_func( @@ -120,7 +121,7 @@ def bremsstrahlung_model_func( core_profiles: state.CoreProfiles, unused_calculated_source_profiles: source_profiles.SourceProfiles | None, unused_conductivity: conductivity_base.Conductivity | None, -) -> tuple[jt.Float[jax.Array, ''], ...]: +) -> tuple[array_typing.Array, ...]: """Model function for the Bremsstrahlung heat sink.""" source_params = runtime_params.sources[source_name] assert isinstance(source_params, RuntimeParams) @@ -142,7 +143,7 @@ class BremsstrahlungHeatSink(source.Source): AFFECTED_CORE_PROFILES: ClassVar[tuple[source.AffectedCoreProfile, ...]] = ( source.AffectedCoreProfile.TEMP_EL, ) - model_func: source.SourceProfileFunction = bremsstrahlung_model_func # pyrefly: ignore[bad-assignment] + model_func: source.SourceProfileFunction = bremsstrahlung_model_func class BremsstrahlungHeatSinkConfig(base.SourceModelBase): @@ -164,7 +165,7 @@ class BremsstrahlungHeatSinkConfig(base.SourceModelBase): @property def model_func(self) -> source.SourceProfileFunction: - return bremsstrahlung_model_func # pyrefly: ignore[bad-return] + return bremsstrahlung_model_func def build_runtime_params( self, diff --git a/torax/_src/sources/cyclotron_radiation_heat_sink.py b/torax/_src/sources/cyclotron_radiation_heat_sink.py index 8bb1821ec..7267683b0 100644 --- a/torax/_src/sources/cyclotron_radiation_heat_sink.py +++ b/torax/_src/sources/cyclotron_radiation_heat_sink.py @@ -16,7 +16,7 @@ """Cyclotron radiation heat sink for electron heat equation..""" import dataclasses -from typing import Annotated, ClassVar, Literal +from typing import Annotated, ClassVar, Literal, Self import chex import jax @@ -34,7 +34,6 @@ from torax._src.sources import source from torax._src.sources import source_profiles from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # Default value for the model function to be used for the Cyclotron radiation # heat sink source. This is also used as an identifier for the model function in @@ -284,7 +283,7 @@ def cyclotron_radiation_albajar( # Dimensionless optical thickness parameter, on-axis: # Simplified form of omega_pe**2 / (c * omega_ce) where omega_pe is the # plasma frequency and omega_ce is the cyclotron frequency. - p_a_0 = 6.04e3 * geo.a_minor * n_e20_face[0] / geo.B_0 # pyrefly: ignore[bad-index] + p_a_0 = 6.04e3 * geo.a_minor * n_e20_face[0] / geo.B_0 # Dimensionless correction term for aspect ratio (equation 15 in Albajar) G = 0.93 * (1 + 0.85 * jnp.exp(-0.82 * geo.R_major_profile / geo.a_minor)) @@ -293,7 +292,7 @@ def cyclotron_radiation_albajar( alpha_n = _alpha_closed_form( beta=2.0, rho_norm=geo.rho_face_norm, - profile_data=n_e20_face, # pyrefly: ignore[bad-argument-type] + profile_data=n_e20_face, profile_edge_value=0.0, ) beta_scan_parameters = ( @@ -303,7 +302,7 @@ def cyclotron_radiation_albajar( ) alpha_t, beta_t = _solve_alpha_t_beta_t_grid_search( rho_norm=geo.rho_face_norm, - te_data=core_profiles.T_e.face_value(), # pyrefly: ignore[bad-argument-type] + te_data=core_profiles.T_e.face_value(), beta_scan_parameters=beta_scan_parameters, # pyrefly: ignore[bad-argument-type] ) @@ -323,10 +322,10 @@ def cyclotron_radiation_albajar( * geo.a_minor**1.38 * geo.elongation_face[-1] ** 0.79 * geo.B_0**2.62 - * n_e20_face[0] ** 0.38 # pyrefly: ignore[bad-index] - * core_profiles.T_e.face_value()[0] # pyrefly: ignore[bad-index] - * (16 + core_profiles.T_e.face_value()[0]) ** 2.61 # pyrefly: ignore[bad-index] - * (1 + 0.12 * core_profiles.T_e.face_value()[0] / p_a_0**0.41) ** -1.51 # pyrefly: ignore[bad-index] + * n_e20_face[0] ** 0.38 + * core_profiles.T_e.face_value()[0] + * (16 + core_profiles.T_e.face_value()[0]) ** 2.61 + * (1 + 0.12 * core_profiles.T_e.face_value()[0] / p_a_0**0.41) ** -1.51 * K * G ) @@ -390,7 +389,7 @@ class CyclotronRadiationHeatSinkConfig(base.SourceModelBase): ) @pydantic.model_validator(mode='after') - def _check_fields(self) -> typing_extensions.Self: + def _check_fields(self) -> Self: if not self.beta_min < self.beta_max: raise ValueError('beta_min must be less than beta_max.') return self diff --git a/torax/_src/sources/electron_cyclotron_source.py b/torax/_src/sources/electron_cyclotron_source.py index 1e3882397..bd0b613c4 100644 --- a/torax/_src/sources/electron_cyclotron_source.py +++ b/torax/_src/sources/electron_cyclotron_source.py @@ -85,9 +85,9 @@ def calc_heating_and_current( ec_power_density = ( source_params.extra_prescribed_power_density + formulas.gaussian_profile( - center=source_params.gaussian_location, # pyrefly: ignore[bad-argument-type] - width=source_params.gaussian_width, # pyrefly: ignore[bad-argument-type] - total=source_params.P_total, # pyrefly: ignore[bad-argument-type] + center=source_params.gaussian_location, + width=source_params.gaussian_width, + total=source_params.P_total, geo=geo, ) ) @@ -131,7 +131,7 @@ class ElectronCyclotronSource(source.Source): source.AffectedCoreProfile.TEMP_EL, source.AffectedCoreProfile.PSI, ) - model_func: source.SourceProfileFunction = calc_heating_and_current # pyrefly: ignore[bad-assignment] + model_func: source.SourceProfileFunction = calc_heating_and_current class ElectronCyclotronSourceConfig(base.SourceModelBase): @@ -177,7 +177,7 @@ class ElectronCyclotronSourceConfig(base.SourceModelBase): @property def model_func(self) -> source.SourceProfileFunction: - return calc_heating_and_current # pyrefly: ignore[bad-return] + return calc_heating_and_current def build_runtime_params( self, diff --git a/torax/_src/sources/formulas.py b/torax/_src/sources/formulas.py index be41ae461..6f67b066b 100644 --- a/torax/_src/sources/formulas.py +++ b/torax/_src/sources/formulas.py @@ -13,8 +13,10 @@ # limitations under the License. """Prescribed formulas for computing source profiles.""" +import chex import jax from jax import numpy as jnp +from torax._src import array_typing from torax._src import math_utils from torax._src.geometry import geometry @@ -24,10 +26,10 @@ def exponential_profile( geo: geometry.Geometry, *, - decay_start: float, - width: float, - total: float, -) -> jax.Array: + decay_start: chex.Numeric, + width: chex.Numeric, + total: chex.Numeric, +) -> array_typing.FloatVectorCell: """Returns an exponential profile on the cell grid. The profile is parameterized by (decay_start, width, total) like so: @@ -50,16 +52,16 @@ def exponential_profile( S = jnp.exp(-(decay_start - r) / width) # calculate constant prefactor C = total / math_utils.volume_integration(S, geo) - return C * S # pyrefly: ignore[bad-return] + return C * S def gaussian_profile( geo: geometry.Geometry, *, - center: float, - width: float, - total: float, -) -> jax.Array: + center: chex.Numeric, + width: chex.Numeric, + total: chex.Numeric, +) -> array_typing.FloatVectorCell: """Returns a gaussian profile on the cell grid. The profile is parameterized by (center, width, total) like so: @@ -82,4 +84,4 @@ def gaussian_profile( S = jnp.exp(-((r - center) ** 2) / (2 * width**2)) # calculate constant prefactor C = total / math_utils.volume_integration(S, geo) - return C * S # pyrefly: ignore[bad-return] + return C * S diff --git a/torax/_src/sources/fusion_heat_source.py b/torax/_src/sources/fusion_heat_source.py index ca3a46990..35818abc4 100644 --- a/torax/_src/sources/fusion_heat_source.py +++ b/torax/_src/sources/fusion_heat_source.py @@ -130,7 +130,7 @@ def calc_fusion( alpha_mass = 4.002602 frac_i = collisions.fast_ion_fractional_heating_formula( birth_energy, - core_profiles.T_e.value, # pyrefly: ignore[bad-argument-type] + core_profiles.T_e.value, alpha_mass, ) frac_e = 1.0 - frac_i @@ -167,7 +167,7 @@ class FusionHeatSource(source.Source): source.AffectedCoreProfile.TEMP_ION, source.AffectedCoreProfile.TEMP_EL, ) - model_func: source.SourceProfileFunction = fusion_heat_model_func # pyrefly: ignore[bad-assignment] + model_func: source.SourceProfileFunction = fusion_heat_model_func class FusionHeatSourceConfig(base.SourceModelBase): @@ -182,7 +182,7 @@ class FusionHeatSourceConfig(base.SourceModelBase): @property def model_func(self) -> source.SourceProfileFunction: - return fusion_heat_model_func # pyrefly: ignore[bad-return] + return fusion_heat_model_func def build_runtime_params( self, diff --git a/torax/_src/sources/gas_puff_source.py b/torax/_src/sources/gas_puff_source.py index 8302163f5..87728bb65 100644 --- a/torax/_src/sources/gas_puff_source.py +++ b/torax/_src/sources/gas_puff_source.py @@ -59,8 +59,8 @@ def calc_puff_source( return ( formulas.exponential_profile( decay_start=1.0, - width=source_params.puff_decay_length, # pyrefly: ignore[bad-argument-type] - total=source_params.S_total, # pyrefly: ignore[bad-argument-type] + width=source_params.puff_decay_length, + total=source_params.S_total, geo=geo, ), ) @@ -74,7 +74,7 @@ class GasPuffSource(source.Source): AFFECTED_CORE_PROFILES: ClassVar[tuple[source.AffectedCoreProfile, ...]] = ( source.AffectedCoreProfile.NE, ) - model_func: source.SourceProfileFunction = calc_puff_source # pyrefly: ignore[bad-assignment] + model_func: source.SourceProfileFunction = calc_puff_source class GasPuffSourceConfig(base.SourceModelBase): @@ -101,7 +101,7 @@ class GasPuffSourceConfig(base.SourceModelBase): @property def model_func(self) -> source.SourceProfileFunction: - return calc_puff_source # pyrefly: ignore[bad-return] + return calc_puff_source def build_runtime_params( self, diff --git a/torax/_src/sources/generic_current_source.py b/torax/_src/sources/generic_current_source.py index 0ba25a1a6..88415501a 100644 --- a/torax/_src/sources/generic_current_source.py +++ b/torax/_src/sources/generic_current_source.py @@ -107,7 +107,7 @@ class GenericCurrentSource(source.Source): AFFECTED_CORE_PROFILES: ClassVar[tuple[source.AffectedCoreProfile, ...]] = ( source.AffectedCoreProfile.PSI, ) - model_func: source.SourceProfileFunction = calculate_generic_current # pyrefly: ignore[bad-assignment] + model_func: source.SourceProfileFunction = calculate_generic_current class GenericCurrentSourceConfig(source_base.SourceModelBase): @@ -145,7 +145,7 @@ class GenericCurrentSourceConfig(source_base.SourceModelBase): @property def model_func(self) -> source.SourceProfileFunction: - return calculate_generic_current # pyrefly: ignore[bad-return] + return calculate_generic_current def build_runtime_params( self, diff --git a/torax/_src/sources/generic_ion_el_heat_source.py b/torax/_src/sources/generic_ion_el_heat_source.py index 4e712f9f8..6ff7e2c9e 100644 --- a/torax/_src/sources/generic_ion_el_heat_source.py +++ b/torax/_src/sources/generic_ion_el_heat_source.py @@ -113,7 +113,7 @@ class GenericIonElectronHeatSource(source.Source): source.AffectedCoreProfile.TEMP_ION, source.AffectedCoreProfile.TEMP_EL, ) - model_func: source.SourceProfileFunction = default_formula # pyrefly: ignore[bad-assignment] + model_func: source.SourceProfileFunction = default_formula class GenericIonElHeatSourceConfig(base.SourceModelBase): @@ -151,7 +151,7 @@ class GenericIonElHeatSourceConfig(base.SourceModelBase): @property def model_func(self) -> source.SourceProfileFunction: - return default_formula # pyrefly: ignore[bad-return] + return default_formula def build_runtime_params( self, diff --git a/torax/_src/sources/generic_particle_source.py b/torax/_src/sources/generic_particle_source.py index d2849d7ca..b7b74756d 100644 --- a/torax/_src/sources/generic_particle_source.py +++ b/torax/_src/sources/generic_particle_source.py @@ -49,9 +49,9 @@ def calc_generic_particle_source( assert isinstance(source_params, RuntimeParams) return ( formulas.gaussian_profile( - center=source_params.deposition_location, # pyrefly: ignore[bad-argument-type] - width=source_params.particle_width, # pyrefly: ignore[bad-argument-type] - total=source_params.S_total, # pyrefly: ignore[bad-argument-type] + center=source_params.deposition_location, + width=source_params.particle_width, + total=source_params.S_total, geo=geo, ), ) @@ -65,7 +65,7 @@ class GenericParticleSource(source.Source): AFFECTED_CORE_PROFILES: ClassVar[tuple[source.AffectedCoreProfile, ...]] = ( source.AffectedCoreProfile.NE, ) - model_func: source.SourceProfileFunction = calc_generic_particle_source # pyrefly: ignore[bad-assignment] + model_func: source.SourceProfileFunction = calc_generic_particle_source @jax.tree_util.register_dataclass @@ -106,7 +106,7 @@ class GenericParticleSourceConfig(base.SourceModelBase): @property def model_func(self) -> source.SourceProfileFunction: - return calc_generic_particle_source # pyrefly: ignore[bad-return] + return calc_generic_particle_source def build_runtime_params( self, diff --git a/torax/_src/sources/impurity_radiation_heat_sink/impurity_radiation_constant_fraction.py b/torax/_src/sources/impurity_radiation_heat_sink/impurity_radiation_constant_fraction.py index 01a011327..b44420bd2 100644 --- a/torax/_src/sources/impurity_radiation_heat_sink/impurity_radiation_constant_fraction.py +++ b/torax/_src/sources/impurity_radiation_heat_sink/impurity_radiation_constant_fraction.py @@ -124,4 +124,4 @@ def build_source( @property def model_func(self) -> source_lib.SourceProfileFunction: - return radially_constant_fraction_of_Pin # pyrefly: ignore[bad-return] + return radially_constant_fraction_of_Pin diff --git a/torax/_src/sources/impurity_radiation_heat_sink/impurity_radiation_mavrin_fit.py b/torax/_src/sources/impurity_radiation_heat_sink/impurity_radiation_mavrin_fit.py index 03aac915c..590899065 100644 --- a/torax/_src/sources/impurity_radiation_heat_sink/impurity_radiation_mavrin_fit.py +++ b/torax/_src/sources/impurity_radiation_heat_sink/impurity_radiation_mavrin_fit.py @@ -144,7 +144,7 @@ class ImpurityRadiationHeatSinkMavrinFitConfig(base.SourceModelBase): @property def model_func(self) -> source_lib.SourceProfileFunction: - return impurity_radiation_mavrin_fit # pyrefly: ignore[bad-return] + return impurity_radiation_mavrin_fit def build_runtime_params( self, diff --git a/torax/_src/sources/ion_cyclotron_source/scaled_profile.py b/torax/_src/sources/ion_cyclotron_source/scaled_profile.py index 1010b191c..496e89e08 100644 --- a/torax/_src/sources/ion_cyclotron_source/scaled_profile.py +++ b/torax/_src/sources/ion_cyclotron_source/scaled_profile.py @@ -169,7 +169,7 @@ class ScaledProfileIonCyclotronSourceConfig(base.IonCyclotronSourceConfig): @property def model_func(self) -> source.SourceProfileFunction: - return scaled_profile_model_func # pyrefly: ignore[bad-return] + return scaled_profile_model_func def build_runtime_params( self, diff --git a/torax/_src/sources/ion_cyclotron_source/toric_nn.py b/torax/_src/sources/ion_cyclotron_source/toric_nn.py index 7b133369e..326796109 100644 --- a/torax/_src/sources/ion_cyclotron_source/toric_nn.py +++ b/torax/_src/sources/ion_cyclotron_source/toric_nn.py @@ -18,7 +18,7 @@ import json import logging import os # pylint: disable=unused-import -from typing import Annotated, Any, Final, Literal, Sequence +from typing import Annotated, Any, Final, Literal, Self, Sequence, cast import chex import flax.linen as nn @@ -43,7 +43,6 @@ from torax._src.sources import source_profiles from torax._src.sources.ion_cyclotron_source import base from torax._src.torax_pydantic import torax_pydantic -import typing_extensions # Internal import. @@ -107,11 +106,11 @@ class ToricNNOutputs: """Outputs from the ToricNN model.""" # Power deposition on helium-3 in MW/m^3/MW_{abs}. - power_deposition_He3: array_typing.FloatVector + power_deposition_He3: array_typing.Array # Power deposition on tritium (second harmonic) in MW/m^3/MW_{abs}. - power_deposition_2T: array_typing.FloatVector + power_deposition_2T: array_typing.Array # Power deposition on electrons in MW/m^3/MW_{abs}. - power_deposition_e: array_typing.FloatVector + power_deposition_e: array_typing.Array class _ToricNN(nn.Module): @@ -264,7 +263,7 @@ def _load_params(self, network_name: str) -> dict[str, Any]: def __hash__(self) -> int: return hash(self._path) - def __eq__(self, other: typing_extensions.Self) -> bool: # pyrefly: ignore[bad-override] + def __eq__(self, other: object) -> bool: return isinstance(other, ToricNNWrapper) @@ -274,7 +273,7 @@ def _toric_nn_predict( inputs: ToricNNInputs, ) -> ToricNNOutputs: """Make a prediction given the inputs.""" - inputs = jnp.array( # pyrefly: ignore[bad-assignment] + inputs_array = jnp.array( [ inputs.frequency, inputs.volume_average_temperature, @@ -289,19 +288,28 @@ def _toric_nn_predict( ], dtype=jax_utils.get_dtype(), ) - outputs_He3 = toric_nn.power_deposition_network.apply( - toric_nn.power_deposition_He3_params, inputs + outputs_He3 = cast( + array_typing.Array, + toric_nn.power_deposition_network.apply( + toric_nn.power_deposition_He3_params, inputs_array + ), ) - outputs_2T = toric_nn.power_deposition_network.apply( - toric_nn.power_deposition_2T_params, inputs + outputs_2T = cast( + array_typing.Array, + toric_nn.power_deposition_network.apply( + toric_nn.power_deposition_2T_params, inputs_array + ), ) - outputs_e = toric_nn.power_deposition_network.apply( - toric_nn.power_deposition_e_params, inputs + outputs_e = cast( + array_typing.Array, + toric_nn.power_deposition_network.apply( + toric_nn.power_deposition_e_params, inputs_array + ), ) return ToricNNOutputs( - power_deposition_He3=outputs_He3, # pyrefly: ignore[bad-argument-type] - power_deposition_2T=outputs_2T, # pyrefly: ignore[bad-argument-type] - power_deposition_e=outputs_e, # pyrefly: ignore[bad-argument-type] + power_deposition_He3=outputs_He3, + power_deposition_2T=outputs_2T, + power_deposition_e=outputs_e, ) @@ -323,7 +331,7 @@ def _get_minority_concentration_from_composition( plasma_composition: plasma_composition_lib.RuntimeParams, core_profiles: state.CoreProfiles, minority_species: str, -) -> jax.Array: +) -> array_typing.Array: """Extract minority species concentration from core profiles. Args: @@ -342,7 +350,7 @@ def _get_minority_concentration_from_composition( if minority_species in plasma_composition.main_ion_names: # For main ions, concentration is fraction * n_i / n_e fraction = core_profiles.main_ion_fractions[minority_species] - return core_profiles.n_i.value * fraction / core_profiles.n_e.value # pyrefly: ignore[bad-return] + return core_profiles.n_i.value * fraction / core_profiles.n_e.value if minority_species in plasma_composition.impurity_names: impurity_fractions = core_profiles.impurity_fractions @@ -351,7 +359,7 @@ def _get_minority_concentration_from_composition( n_imp_species = ( fraction * core_profiles.n_impurity.value * impurity_density_scaling ) - return n_imp_species / core_profiles.n_e.value # pyrefly: ignore[bad-return] + return n_imp_species / core_profiles.n_e.value raise ValueError( f'Minority species {minority_species} not found in plasma composition.' @@ -393,23 +401,24 @@ def icrh_model_func( else: # Use legacy parameter (backward compatibility) # TODO(b/434175938): Remove backward compatibility in V2. + assert source_params.minority_concentration is not None minority_concentration_scalar = source_params.minority_concentration # For profile-dependent calculations, use constant value minority_concentration_profile = source_params.minority_concentration # Construct inputs for ToricNN. volume_average_temperature = math_utils.volume_average( - core_profiles.T_e.value, geo # pyrefly: ignore[bad-argument-type] + core_profiles.T_e.value, geo ) volume_average_density = math_utils.volume_average( - core_profiles.n_e.value, geo # pyrefly: ignore[bad-argument-type] + core_profiles.n_e.value, geo ) # Peaking factors are core w.r.t volume averages. temperature_peaking_factor = ( - core_profiles.T_e.value[0] / volume_average_temperature # pyrefly: ignore[bad-index] + core_profiles.T_e.value[0] / volume_average_temperature ) - density_peaking_factor = core_profiles.n_e.value[0] / volume_average_density # pyrefly: ignore[bad-index] + density_peaking_factor = core_profiles.n_e.value[0] / volume_average_density Router = geo.R_out_face[-1] # Use LCFS outboard radius Rinner = geo.R_in_face[-1] # Use LCFS inboard radius # Assumption: inner and outer gaps are not functions of z0. @@ -422,11 +431,11 @@ def icrh_model_func( volume_average_temperature=volume_average_temperature, volume_average_density=volume_average_density / 1e20, # convert to 10^20 m^-3 - minority_concentration=minority_concentration_scalar # pyrefly: ignore[unsupported-operation] + minority_concentration=minority_concentration_scalar * 100, # Convert to percentage. gap_inner=gap_inner, gap_outer=gap_outer, - z0=geo.z_magnetic_axis(), # pyrefly: ignore[bad-argument-type] + z0=geo.z_magnetic_axis(), temperature_peaking_factor=temperature_peaking_factor, density_peaking_factor=density_peaking_factor, B_0=geo.B_0, @@ -468,19 +477,19 @@ def icrh_model_func( n_tail, T_tail = fast_ion_utils.bimaxwellian_split( power_deposition=power_deposition_he3, - T_e=core_profiles.T_e.value, # pyrefly: ignore[bad-argument-type] - n_e=core_profiles.n_e.value, # pyrefly: ignore[bad-argument-type] - T_i=core_profiles.T_i.value, # pyrefly: ignore[bad-argument-type] - n_i=core_profiles.n_i.value, # pyrefly: ignore[bad-argument-type] - minority_concentration=minority_concentration_profile, # pyrefly: ignore[bad-argument-type] - P_total_W=source_params.P_total, # pyrefly: ignore[bad-argument-type] + T_e=core_profiles.T_e.value, + n_e=core_profiles.n_e.value, + T_i=core_profiles.T_i.value, + n_i=core_profiles.n_i.value, + minority_concentration=minority_concentration_profile, + P_total_W=source_params.P_total, charge_number=he3_charge_number, mass_number=he3_atomic_mass, - bulk_ion_mass=core_profiles.A_i, # pyrefly: ignore[bad-argument-type] - Z_i=core_profiles.Z_i, # pyrefly: ignore[bad-argument-type] - n_impurity=core_profiles.n_impurity.value, # pyrefly: ignore[bad-argument-type] - Z_impurity=core_profiles.Z_impurity, # pyrefly: ignore[bad-argument-type] - A_impurity=core_profiles.A_impurity, # pyrefly: ignore[bad-argument-type] + bulk_ion_mass=core_profiles.A_i, + Z_i=core_profiles.Z_i, + n_impurity=core_profiles.n_impurity.value, + Z_impurity=core_profiles.Z_impurity, + A_impurity=core_profiles.A_impurity, ) # Build fast ion output for all supported species. @@ -509,7 +518,7 @@ def icrh_model_func( frac_ion_heating = collisions.fast_ion_fractional_heating_formula( T_tail, - core_profiles.T_e.value, # pyrefly: ignore[bad-argument-type] + core_profiles.T_e.value, he3_atomic_mass, ) absorbed_power = source_params.P_total * source_params.absorption_fraction @@ -592,7 +601,7 @@ def build_source(self) -> base.IonCyclotronSource: return base.IonCyclotronSource(model_func=None) @pydantic.model_validator(mode='after') - def _validate_minority_species(self) -> typing_extensions.Self: + def _validate_minority_species(self) -> Self: if self.minority_species is not None and self.minority_species != 'He3': raise ValueError( "Minority species must be 'He3' if specified. Got:" @@ -604,7 +613,7 @@ def _validate_minority_species(self) -> typing_extensions.Self: @pydantic.model_validator(mode='after') def _log_warning_for_used_minority_concentration( self, - ) -> typing_extensions.Self: + ) -> Self: """Logs a warning if minority_concentration is provided.""" if self.minority_concentration is not None: logging.warning( diff --git a/torax/_src/sources/ohmic_heat_source.py b/torax/_src/sources/ohmic_heat_source.py index 2dcd7681f..c8ae89d2c 100644 --- a/torax/_src/sources/ohmic_heat_source.py +++ b/torax/_src/sources/ohmic_heat_source.py @@ -73,7 +73,7 @@ def ohmic_model_func( psidot = psi_calculations.calculate_psidot_from_psi_sources( psi_sources=psi_sources, sigma=conductivity.sigma, - resistivity_multiplier=runtime_params.numerics.resistivity_multiplier, # pyrefly: ignore[bad-argument-type] + resistivity_multiplier=runtime_params.numerics.resistivity_multiplier, psi=core_profiles.psi, geo=geo, ) @@ -96,7 +96,7 @@ class OhmicHeatSource(source_lib.Source): AFFECTED_CORE_PROFILES: ClassVar[ tuple[source_lib.AffectedCoreProfile, ...] ] = (source_lib.AffectedCoreProfile.TEMP_EL,) - model_func: source_lib.SourceProfileFunction = ohmic_model_func # pyrefly: ignore[bad-assignment] + model_func: source_lib.SourceProfileFunction = ohmic_model_func class OhmicHeatSourceConfig(base.SourceModelBase): @@ -111,7 +111,7 @@ class OhmicHeatSourceConfig(base.SourceModelBase): @property def model_func(self) -> source_lib.SourceProfileFunction: - return ohmic_model_func # pyrefly: ignore[bad-return] + return ohmic_model_func def build_runtime_params( self, diff --git a/torax/_src/sources/pellet_source.py b/torax/_src/sources/pellet_source.py index 465f7f85d..5edcd6d55 100644 --- a/torax/_src/sources/pellet_source.py +++ b/torax/_src/sources/pellet_source.py @@ -48,9 +48,9 @@ def calc_pellet_source( assert isinstance(source_params, RuntimeParams) return ( formulas.gaussian_profile( - center=source_params.pellet_deposition_location, # pyrefly: ignore[bad-argument-type] - width=source_params.pellet_width, # pyrefly: ignore[bad-argument-type] - total=source_params.S_total, # pyrefly: ignore[bad-argument-type] + center=source_params.pellet_deposition_location, + width=source_params.pellet_width, + total=source_params.S_total, geo=geo, ), ) @@ -64,7 +64,7 @@ class PelletSource(source.Source): AFFECTED_CORE_PROFILES: ClassVar[tuple[source.AffectedCoreProfile, ...]] = ( source.AffectedCoreProfile.NE, ) - model_func: source.SourceProfileFunction = calc_pellet_source # pyrefly: ignore[bad-assignment] + model_func: source.SourceProfileFunction = calc_pellet_source @jax.tree_util.register_dataclass @@ -105,7 +105,7 @@ class PelletSourceConfig(base.SourceModelBase): @property def model_func(self) -> source.SourceProfileFunction: - return calc_pellet_source # pyrefly: ignore[bad-return] + return calc_pellet_source def build_runtime_params( self, diff --git a/torax/_src/sources/pydantic_model.py b/torax/_src/sources/pydantic_model.py index 7077e9d78..f3a1df2e3 100644 --- a/torax/_src/sources/pydantic_model.py +++ b/torax/_src/sources/pydantic_model.py @@ -15,7 +15,7 @@ """Pydantic config for source models.""" import copy -from typing import Any +from typing import Any, Self import immutabledict import pydantic @@ -39,7 +39,6 @@ from torax._src.sources.ion_cyclotron_source import scaled_profile from torax._src.sources.ion_cyclotron_source import toric_nn from torax._src.torax_pydantic import torax_pydantic -from typing_extensions import Self class Sources(torax_pydantic.BaseModelFrozen): diff --git a/torax/_src/sources/source.py b/torax/_src/sources/source.py index 0ec6f4c15..f68a9d054 100644 --- a/torax/_src/sources/source.py +++ b/torax/_src/sources/source.py @@ -54,6 +54,7 @@ def __call__( core_profiles: state.CoreProfiles, calculated_source_profiles: source_profiles.SourceProfiles | None, unused_conductivity: conductivity_base.Conductivity | None, + /, ) -> tuple[SourceProfileElement, ...]: ... diff --git a/torax/_src/sources/source_profiles.py b/torax/_src/sources/source_profiles.py index b5d584a48..b61a7b443 100644 --- a/torax/_src/sources/source_profiles.py +++ b/torax/_src/sources/source_profiles.py @@ -16,7 +16,7 @@ from collections.abc import Iterator, Mapping import dataclasses import operator -from typing import Literal +from typing import Any, Literal, Self import jax import jax.numpy as jnp @@ -27,7 +27,6 @@ from torax._src.output_tools import output_grid_context from torax._src.output_tools import output_keys from torax._src.physics import fast_ion as fast_ion_lib -import typing_extensions # pylint: disable=invalid-name @@ -46,7 +45,7 @@ class QeiInfo: p_ei: array_typing.Array @classmethod - def zeros(cls, geo: geometry.Geometry) -> typing_extensions.Self: + def zeros(cls, geo: geometry.Geometry) -> Self: return cls( implicit_ii=jnp.zeros_like(geo.rho), explicit_i=jnp.zeros_like(geo.rho), @@ -111,9 +110,9 @@ class SourceProfiles: @classmethod def merge( cls, - explicit_source_profiles: typing_extensions.Self, - implicit_source_profiles: typing_extensions.Self, - ) -> typing_extensions.Self: + explicit_source_profiles: Self, + implicit_source_profiles: Self, + ) -> Self: """Returns a SourceProfiles that merges the input profiles. Sources can either be explicit or implicit. The explicit_source_profiles @@ -139,14 +138,14 @@ def merge( implicit (assuming the source model outputted a non-zero profile). """ - def _is_fast_ions_dict(x: typing_extensions.Any) -> bool: + def _is_fast_ions_dict(x: Any) -> bool: return isinstance(x, dict) and all( isinstance(v, tuple) and all(isinstance(el, fast_ion_lib.FastIon) for el in v) for v in x.values() ) - def _merge(a: typing_extensions.Any, b: typing_extensions.Any): + def _merge(a: Any, b: Any): if _is_fast_ions_dict(a): return {**a, **b} return operator.add(a, b) diff --git a/torax/_src/sources/tests/constant_fraction_impurity_radiation_heat_sink_test.py b/torax/_src/sources/tests/constant_fraction_impurity_radiation_heat_sink_test.py index 9f1171db4..f6f71250c 100644 --- a/torax/_src/sources/tests/constant_fraction_impurity_radiation_heat_sink_test.py +++ b/torax/_src/sources/tests/constant_fraction_impurity_radiation_heat_sink_test.py @@ -76,7 +76,7 @@ def test_source_value(self): ) heat_source = generic_ion_el_heat_source.GenericIonElectronHeatSource( - model_func=generic_ion_el_heat_source.default_formula, # pyrefly: ignore[bad-argument-type] + model_func=generic_ion_el_heat_source.default_formula, ) geo = circular_geometry.CircularConfig().build_geometry() @@ -89,7 +89,7 @@ def test_source_value(self): ) impurity_radiation_sink = impurity_radiation_heat_sink_lib.ImpurityRadiationHeatSink( - model_func=impurity_radiation_constant_fraction.radially_constant_fraction_of_Pin # pyrefly: ignore[bad-argument-type] + model_func=impurity_radiation_constant_fraction.radially_constant_fraction_of_Pin ) impurity_radiation_heat_sink_power_density = ( @@ -100,8 +100,8 @@ def test_source_value(self): calculated_source_profiles=source_profiles.SourceProfiles( bootstrap_current=mock.ANY, qei=mock.ANY, - T_e={'foo': el}, # pyrefly: ignore[bad-argument-type, bad-assignment] - T_i={'foo_source': ion}, # pyrefly: ignore[bad-argument-type, bad-assignment] + T_e={'foo': el}, # pyrefly: ignore[bad-argument-type] + T_i={'foo_source': ion}, # pyrefly: ignore[bad-argument-type] ), conductivity=None, ) diff --git a/torax/_src/sources/tests/ohmic_heat_source_test.py b/torax/_src/sources/tests/ohmic_heat_source_test.py index 6c429aa6b..60b0a8d85 100644 --- a/torax/_src/sources/tests/ohmic_heat_source_test.py +++ b/torax/_src/sources/tests/ohmic_heat_source_test.py @@ -32,7 +32,7 @@ class OhmicHeatSourceTest(test_lib.SingleProfileSourceTestCase): def test_raises_error_if_calculated_source_profiles_is_none(self): source = ohmic_heat_source.OhmicHeatSource( - model_func=ohmic_heat_source.ohmic_model_func # pyrefly: ignore[bad-argument-type] + model_func=ohmic_heat_source.ohmic_model_func ) source_config = self._source_config_class.from_dict({}) face_centers = interpolated_param_2d.get_face_centers(4) @@ -61,7 +61,7 @@ def test_raises_error_if_calculated_source_profiles_is_none(self): def test_raises_error_if_conductivity_is_none(self): source = ohmic_heat_source.OhmicHeatSource( - model_func=ohmic_heat_source.ohmic_model_func # pyrefly: ignore[bad-argument-type] + model_func=ohmic_heat_source.ohmic_model_func ) source_config = self._source_config_class.from_dict({}) face_centers = interpolated_param_2d.get_face_centers(4) diff --git a/torax/_src/sources/tests/register_model_test.py b/torax/_src/sources/tests/register_model_test.py index 1110c8e15..d7cc1e187 100644 --- a/torax/_src/sources/tests/register_model_test.py +++ b/torax/_src/sources/tests/register_model_test.py @@ -70,7 +70,7 @@ class NewGasPuffSourceModelConfig(source_base_pydantic_model.SourceModelBase): @property def model_func(self) -> source_lib.SourceProfileFunction: - return double_gas_puff_source # pyrefly: ignore[bad-return] + return double_gas_puff_source def build_source(self) -> source_lib.Source: return gas_puff_source_lib.GasPuffSource(model_func=self.model_func) @@ -100,7 +100,7 @@ class DuplicateGasPuffSourceModelConfig( @property def model_func(self) -> source_lib.SourceProfileFunction: - return double_gas_puff_source # pyrefly: ignore[bad-return] + return double_gas_puff_source def build_source(self) -> source_lib.Source: return gas_puff_source_lib.GasPuffSource(model_func=self.model_func) diff --git a/torax/_src/sources/tests/source_profile_builders_test.py b/torax/_src/sources/tests/source_profile_builders_test.py index 1c5e047c2..37263fb86 100644 --- a/torax/_src/sources/tests/source_profile_builders_test.py +++ b/torax/_src/sources/tests/source_profile_builders_test.py @@ -86,7 +86,7 @@ class TestSource(source.Source): AFFECTED_CORE_PROFILES = (source.AffectedCoreProfile.PSI,) test_source = TestSource( - model_func=lambda *args: (jnp.ones(self.geo.rho.shape),) # pyrefly: ignore[bad-argument-type] + model_func=lambda *args: (jnp.ones(self.geo.rho.shape),) ) source_models = mock.create_autospec( source_models_lib.SourceModels, @@ -136,7 +136,7 @@ class TestSource(source.Source): ) test_source = TestSource( - model_func=lambda *args: (jnp.ones_like(self.geo.rho),) * 2 # pyrefly: ignore[bad-argument-type] + model_func=lambda *args: (jnp.ones_like(self.geo.rho),) * 2 ) source_models = mock.create_autospec( source_models_lib.SourceModels, @@ -216,7 +216,7 @@ class TestSource(source.Source): AFFECTED_CORE_PROFILES = (source.AffectedCoreProfile.PSI,) test_source = TestSource( - model_func=lambda *args: (jnp.ones(self.geo.rho.shape),) # pyrefly: ignore[bad-argument-type] + model_func=lambda *args: (jnp.ones(self.geo.rho.shape),) ) source_models = mock.create_autospec( source_models_lib.SourceModels, diff --git a/torax/_src/state.py b/torax/_src/state.py index 551728464..d0ebe768f 100644 --- a/torax/_src/state.py +++ b/torax/_src/state.py @@ -17,7 +17,7 @@ import dataclasses import enum import functools -from typing import Mapping +from typing import Mapping, Self from absl import logging import jax @@ -31,7 +31,6 @@ from torax._src.output_tools import output_keys from torax._src.physics import charge_states from torax._src.physics import fast_ion as fast_ion_lib -import typing_extensions # pylint: disable=invalid-name @@ -193,12 +192,12 @@ def n_impurity_thermal(self) -> cell_variable.CellVariable: n_impurity_thermal_right = self.n_impurity.right_face_constraint for fast_ion in self.fast_ions: if fast_ion.species in self.impurity_fractions: - n_impurity_thermal_value -= fast_ion.n.value # pyrefly: ignore[unsupported-operation] + n_impurity_thermal_value -= fast_ion.n.value if ( n_impurity_thermal_right is not None and fast_ion.n.right_face_constraint is not None ): - n_impurity_thermal_right -= fast_ion.n.right_face_constraint # pyrefly: ignore[unsupported-operation] + n_impurity_thermal_right -= fast_ion.n.right_face_constraint return cell_variable.CellVariable( value=n_impurity_thermal_value, face_centers=self.n_impurity.face_centers, @@ -299,7 +298,7 @@ def quasineutrality_satisfied(self) -> bool: self.n_e.value, ).item() - def negative_temperature_or_density(self) -> jax.Array: + def negative_temperature_or_density(self) -> bool: """Checks if any temperature or density is negative.""" profiles_to_check = ( self.T_i, @@ -311,11 +310,10 @@ def negative_temperature_or_density(self) -> jax.Array: ) # Check if any profile is less than -eps # (allowing for numerical precision errors) - return np.any( # pyrefly: ignore[bad-return] - np.array([ - np.any(np.less(x, -constants.CONSTANTS.eps)) # pyrefly: ignore[unsupported-operation] - for x in jax.tree.leaves(profiles_to_check) - ]) + eps = float(constants.CONSTANTS.eps) + return any( + bool(np.any(x < -eps)) + for x in jax.tree.leaves(profiles_to_check) ) def below_minimum_temperature(self, T_minimum_eV: float) -> bool: @@ -481,34 +479,34 @@ class CoreTransport: `transport_model/transport_model.py` for more details. """ - chi_face_ion: jax.Array - chi_face_el: jax.Array - d_face_el: jax.Array - v_face_el: jax.Array - chi_face_el_bohm: jax.Array | None = None - chi_face_el_gyrobohm: jax.Array | None = None - chi_face_ion_bohm: jax.Array | None = None - chi_face_ion_gyrobohm: jax.Array | None = None - chi_face_el_itg: jax.Array | None = None - chi_face_el_tem: jax.Array | None = None - chi_face_el_etg: jax.Array | None = None - chi_face_ion_itg: jax.Array | None = None - chi_face_ion_tem: jax.Array | None = None - d_face_el_itg: jax.Array | None = None - d_face_el_tem: jax.Array | None = None - v_face_el_itg: jax.Array | None = None - v_face_el_tem: jax.Array | None = None - chi_neo_i: jax.Array | None = None - chi_neo_e: jax.Array | None = None - D_neo_e: jax.Array | None = None - V_neo_e: jax.Array | None = None - V_neo_ware_e: jax.Array | None = None - chi_face_ion_pereverzev: jax.Array | None = None - chi_face_el_pereverzev: jax.Array | None = None - full_v_heat_face_ion_pereverzev: jax.Array | None = None - full_v_heat_face_el_pereverzev: jax.Array | None = None - d_face_el_pereverzev: jax.Array | None = None - v_face_el_pereverzev: jax.Array | None = None + chi_face_ion: array_typing.FloatVectorFace + chi_face_el: array_typing.FloatVectorFace + d_face_el: array_typing.FloatVectorFace + v_face_el: array_typing.FloatVectorFace + chi_face_el_bohm: array_typing.FloatVectorFace | None = None + chi_face_el_gyrobohm: array_typing.FloatVectorFace | None = None + chi_face_ion_bohm: array_typing.FloatVectorFace | None = None + chi_face_ion_gyrobohm: array_typing.FloatVectorFace | None = None + chi_face_el_itg: array_typing.FloatVectorFace | None = None + chi_face_el_tem: array_typing.FloatVectorFace | None = None + chi_face_el_etg: array_typing.FloatVectorFace | None = None + chi_face_ion_itg: array_typing.FloatVectorFace | None = None + chi_face_ion_tem: array_typing.FloatVectorFace | None = None + d_face_el_itg: array_typing.FloatVectorFace | None = None + d_face_el_tem: array_typing.FloatVectorFace | None = None + v_face_el_itg: array_typing.FloatVectorFace | None = None + v_face_el_tem: array_typing.FloatVectorFace | None = None + chi_neo_i: array_typing.FloatVectorFace | None = None + chi_neo_e: array_typing.FloatVectorFace | None = None + D_neo_e: array_typing.FloatVectorFace | None = None + V_neo_e: array_typing.FloatVectorFace | None = None + V_neo_ware_e: array_typing.FloatVectorFace | None = None + chi_face_ion_pereverzev: array_typing.FloatVectorFace | None = None + chi_face_el_pereverzev: array_typing.FloatVectorFace | None = None + full_v_heat_face_ion_pereverzev: array_typing.FloatVectorFace | None = None + full_v_heat_face_el_pereverzev: array_typing.FloatVectorFace | None = None + d_face_el_pereverzev: array_typing.FloatVectorFace | None = None + v_face_el_pereverzev: array_typing.FloatVectorFace | None = None def __post_init__(self): # Use the array size of chi_face_el as a template. @@ -537,25 +535,34 @@ def __post_init__(self): self.v_face_el_pereverzev = jnp.zeros_like(template) @property - def chi_face_ion_total(self) -> jax.Array: + def chi_face_ion_total(self) -> array_typing.FloatVectorFace: """Calculates the total ion heat diffusion coefficient.""" - return self.chi_face_ion + self.chi_face_ion_pereverzev + self.chi_neo_i # pyrefly: ignore[unsupported-operation] + assert self.chi_face_ion_pereverzev is not None + assert self.chi_neo_i is not None + return self.chi_face_ion + self.chi_face_ion_pereverzev + self.chi_neo_i @property - def chi_face_el_total(self) -> jax.Array: + def chi_face_el_total(self) -> array_typing.FloatVectorFace: """Calculates the total electron heat diffusion coefficient.""" - return self.chi_face_el + self.chi_face_el_pereverzev + self.chi_neo_e # pyrefly: ignore[unsupported-operation] + assert self.chi_face_el_pereverzev is not None + assert self.chi_neo_e is not None + return self.chi_face_el + self.chi_face_el_pereverzev + self.chi_neo_e @property - def d_face_el_total(self) -> jax.Array: + def d_face_el_total(self) -> array_typing.FloatVectorFace: """Calculates the total particle diffusion coefficient.""" - return self.d_face_el + self.d_face_el_pereverzev + self.D_neo_e # pyrefly: ignore[unsupported-operation] + assert self.d_face_el_pereverzev is not None + assert self.D_neo_e is not None + return self.d_face_el + self.d_face_el_pereverzev + self.D_neo_e @property - def v_face_el_total(self) -> jax.Array: + def v_face_el_total(self) -> array_typing.FloatVectorFace: """Calculates the total particle convection coefficient.""" + assert self.v_face_el_pereverzev is not None + assert self.V_neo_e is not None + assert self.V_neo_ware_e is not None return ( - self.v_face_el # pyrefly: ignore[unsupported-operation] + self.v_face_el + self.v_face_el_pereverzev + self.V_neo_e + self.V_neo_ware_e @@ -564,7 +571,7 @@ def v_face_el_total(self) -> jax.Array: def chi_max( self, geo: geometry.Geometry, - ) -> jax.Array: + ) -> array_typing.FloatScalar: """Calculates the maximum value of chi. Args: @@ -573,13 +580,15 @@ def chi_max( Returns: chi_max: Maximum value of chi. """ + assert self.chi_neo_i is not None + assert self.chi_neo_e is not None return jnp.maximum( - jnp.max((self.chi_face_ion + self.chi_neo_i) * geo.g1_over_vpr2_face), # pyrefly: ignore[unsupported-operation] - jnp.max((self.chi_face_el + self.chi_neo_e) * geo.g1_over_vpr2_face), # pyrefly: ignore[unsupported-operation] + jnp.max((self.chi_face_ion + self.chi_neo_i) * geo.g1_over_vpr2_face), + jnp.max((self.chi_face_el + self.chi_neo_e) * geo.g1_over_vpr2_face), ) @classmethod - def zeros(cls, geo: geometry.Geometry) -> typing_extensions.Self: + def zeros(cls, geo: geometry.Geometry) -> Self: """Returns a CoreTransport with all zeros. Useful for initializing.""" shape = geo.rho_face.shape return cls( diff --git a/torax/_src/torax_pydantic/interpolated_param_1d.py b/torax/_src/torax_pydantic/interpolated_param_1d.py index fee13a0b4..627510424 100644 --- a/torax/_src/torax_pydantic/interpolated_param_1d.py +++ b/torax/_src/torax_pydantic/interpolated_param_1d.py @@ -16,7 +16,7 @@ import dataclasses import functools -from typing import Any, TypeAlias +from typing import Annotated, Any, Self, TypeAlias import chex import equinox as eqx @@ -29,7 +29,6 @@ from torax._src.torax_pydantic import interpolated_param_2d from torax._src.torax_pydantic import model_base from torax._src.torax_pydantic import pydantic_types -import typing_extensions @jax.tree_util.register_dataclass @@ -59,10 +58,10 @@ class TimeVaryingScalar(model_base.BaseModelFrozen): time: pydantic_types.NumpyArray1DSorted value: pydantic_types.NumpyArray - is_bool_param: typing_extensions.Annotated[bool, model_base.JAX_STATIC] = ( + is_bool_param: Annotated[bool, model_base.JAX_STATIC] = ( False ) - interpolation_mode: typing_extensions.Annotated[ + interpolation_mode: Annotated[ interpolated_param.InterpolationMode, model_base.JAX_STATIC ] = interpolated_param.InterpolationMode.PIECEWISE_LINEAR @@ -95,7 +94,7 @@ def to_time_varying_array(self) -> interpolated_param_2d.TimeVaryingArray: def update( self, replacements: TimeVaryingScalarUpdate - ) -> typing_extensions.Self: + ) -> Self: """This method can be used under `jax.jit`.""" value = replacements.value if replacements.value is not None else self.value time = replacements.time if replacements.time is not None else self.time @@ -105,7 +104,7 @@ def update( f' be the same length. Got value: {value.shape}, time: {time.shape}.' ) - def get_leaves(x: typing_extensions.Self) -> tuple[chex.Array, chex.Array]: + def get_leaves(x: Self) -> tuple[chex.Array, chex.Array]: return (x.time, x.value) return eqx.tree_at(get_leaves, self, (time, value),) @@ -167,7 +166,7 @@ def __eq__(self, other): ) @pydantic.model_validator(mode='after') - def _ensure_consistent_arrays(self) -> typing_extensions.Self: + def _ensure_consistent_arrays(self) -> Self: if not np.issubdtype(self.time.dtype, np.floating): raise ValueError('The time array must be a float array.') @@ -228,7 +227,7 @@ def _get_cached_interpolated_param( class TimeVaryingScalarStep(TimeVaryingScalar): """TimeVaryingScalar with STEP interpolation mode by default.""" - interpolation_mode: typing_extensions.Annotated[ + interpolation_mode: Annotated[ interpolated_param.InterpolationMode, model_base.JAX_STATIC ] = interpolated_param.InterpolationMode.STEP @@ -280,15 +279,15 @@ def scalar_bounds_validator( ) -PositiveTimeVaryingScalar: TypeAlias = typing_extensions.Annotated[ +PositiveTimeVaryingScalar: TypeAlias = Annotated[ TimeVaryingScalar, scalar_bounds_validator(gt=0.0) ] -NonNegativeTimeVaryingScalar: TypeAlias = typing_extensions.Annotated[ +NonNegativeTimeVaryingScalar: TypeAlias = Annotated[ TimeVaryingScalar, scalar_bounds_validator(ge=0.0) ] -NonNegativeTimeVaryingScalarStep: TypeAlias = typing_extensions.Annotated[ +NonNegativeTimeVaryingScalarStep: TypeAlias = Annotated[ TimeVaryingScalarStep, scalar_bounds_validator(ge=0.0) ] -UnitIntervalTimeVaryingScalar: TypeAlias = typing_extensions.Annotated[ +UnitIntervalTimeVaryingScalar: TypeAlias = Annotated[ TimeVaryingScalar, scalar_bounds_validator(ge=0.0, le=1.0) ] diff --git a/torax/_src/torax_pydantic/interpolated_param_2d.py b/torax/_src/torax_pydantic/interpolated_param_2d.py index 02f42a525..06107ed24 100644 --- a/torax/_src/torax_pydantic/interpolated_param_2d.py +++ b/torax/_src/torax_pydantic/interpolated_param_2d.py @@ -17,7 +17,7 @@ from collections.abc import Mapping import dataclasses import functools -from typing import Any, Literal, TypeAlias +from typing import Annotated, Any, Literal, Self, TypeAlias import chex import equinox as eqx @@ -31,7 +31,6 @@ from torax._src import jax_utils from torax._src.torax_pydantic import model_base from torax._src.torax_pydantic import pydantic_types -import typing_extensions import xarray as xr ValueType: TypeAlias = dict[ @@ -90,7 +89,7 @@ def cell_widths(self) -> jax.Array: """Widths of cells.""" return jnp.diff(self.face_centers) - def __eq__(self, other: typing_extensions.Self) -> bool: # pyrefly: ignore[bad-override] + def __eq__(self, other: object) -> bool: """Custom equality to handle numpy array comparison.""" if not isinstance(other, Grid1D): return False @@ -156,10 +155,10 @@ class TimeVaryingArray(model_base.BaseModelFrozen): """ value: ValueType - rho_interpolation_mode: typing_extensions.Annotated[ + rho_interpolation_mode: Annotated[ interpolated_param.InterpolationMode, model_base.JAX_STATIC ] = interpolated_param.InterpolationMode.PIECEWISE_LINEAR - time_interpolation_mode: typing_extensions.Annotated[ + time_interpolation_mode: Annotated[ interpolated_param.InterpolationMode, model_base.JAX_STATIC ] = interpolated_param.InterpolationMode.PIECEWISE_LINEAR grid: Grid1D | None = None @@ -249,8 +248,8 @@ def _linear_nonpositive_subintervals( if not np.any(sign_change): return subintervals - t_cross = t_i + dt * (-v_i[sign_change]) / ( # pyrefly: ignore[bad-index] - v_next[sign_change] - v_i[sign_change] # pyrefly: ignore[bad-index] + t_cross = t_i + dt * (-v_i[sign_change]) / ( + v_next[sign_change] - v_i[sign_change] ) # Points nonpositive at t_i are nonpositive on [t_i, t_cross]. @@ -360,7 +359,7 @@ def get_value( def update( self, replace_value: TimeVaryingArrayUpdate - ) -> typing_extensions.Self: + ) -> Self: """This method can be used under `jax.jit`.""" assert self.grid is not None, 'grid must be set to update.' @@ -394,7 +393,7 @@ def update( ) def get_leaves( - x: typing_extensions.Self, + x: Self, ) -> tuple[ chex.Array, chex.Array, chex.Array, chex.Array, chex.Array, chex.Array ]: @@ -415,7 +414,9 @@ def get_leaves( (time, cell_value, time, face_value, time, face_right_value), ) - def __eq__(self, other: typing_extensions.Self): # pyrefly: ignore[bad-override] + def __eq__(self, other: object) -> bool: + if not isinstance(other, TimeVaryingArray): + return False try: chex.assert_trees_all_equal(self.value, other.value) return ( @@ -755,7 +756,7 @@ def array_bounds_validator( ) -PositiveTimeVaryingArray: TypeAlias = typing_extensions.Annotated[ +PositiveTimeVaryingArray: TypeAlias = Annotated[ TimeVaryingArray, array_bounds_validator(gt=0.0) ] @@ -905,6 +906,6 @@ def get_face_centers(nx: int, dx: float | None = None) -> np.ndarray: return np.linspace(0, nx * dx, nx + 1) -NonNegativeTimeVaryingArray: TypeAlias = typing_extensions.Annotated[ +NonNegativeTimeVaryingArray: TypeAlias = Annotated[ TimeVaryingArray, array_bounds_validator(ge=0.0) ] diff --git a/torax/_src/torax_pydantic/model_base.py b/torax/_src/torax_pydantic/model_base.py index 55521dd6e..a39886eb3 100644 --- a/torax/_src/torax_pydantic/model_base.py +++ b/torax/_src/torax_pydantic/model_base.py @@ -17,12 +17,11 @@ from collections.abc import Set import functools import inspect -from typing import Any, Final, Mapping, Sequence, TypeAlias +from typing import Any, Final, Mapping, Self, Sequence, TypeAlias import jax import pydantic import treelib -from typing_extensions import Self TIME_INVARIANT: Final[str] = '_pydantic_time_invariant_field' JAX_STATIC: Final[str] = '_pydantic_jax_static_field' @@ -140,7 +139,7 @@ def time_invariant_fields(cls) -> tuple[str, ...]: ) @property - def _direct_submodels(self) -> tuple[Self, ...]: + def _direct_submodels(self) -> tuple['BaseModelFrozen', ...]: """Direct submodels in the model.""" def is_leaf(x): @@ -153,10 +152,10 @@ def is_leaf(x): # Some Pydantic models are values of a dict. We flatten the tree to access # them. leaves = jax.tree.flatten(leaves, is_leaf=is_leaf)[0] - return tuple(i for i in leaves if isinstance(i, BaseModelFrozen)) # pyrefly: ignore[bad-return] + return tuple(i for i in leaves if isinstance(i, BaseModelFrozen)) @property - def submodels(self) -> tuple[Self, ...]: + def submodels(self) -> tuple['BaseModelFrozen', ...]: """A tuple of the model and all submodels. This will return all Pydantic models directly inside model fields, and @@ -166,7 +165,7 @@ def submodels(self) -> tuple[Self, ...]: A tuple of the model and all model submodels. """ - all_submodels = [self] + all_submodels: list[BaseModelFrozen] = [self] new_submodels = self._direct_submodels while new_submodels: new_submodels_temp = [] @@ -286,7 +285,7 @@ def _update_fields(self, x: Mapping[str, Any]): # Re-validate all ancestral models. m.__class__.from_dict(m.to_dict()) - def _lookup_path(self, paths: Sequence[str]) -> Self: + def _lookup_path(self, paths: Sequence[str]) -> 'BaseModelFrozen': """Returns the model at the given path.""" value = self for path in paths: @@ -303,4 +302,4 @@ def _lookup_path(self, paths: Sequence[str]) -> Self: raise ValueError(f'Cannot look up path {path} in {value}') if not isinstance(value, BaseModelFrozen): raise ValueError(f'The value at path {paths} is not a Pydantic model.') - return value # pyrefly: ignore[bad-return] + return value diff --git a/torax/_src/torax_pydantic/model_config.py b/torax/_src/torax_pydantic/model_config.py index 658062ac1..1ef9bca11 100644 --- a/torax/_src/torax_pydantic/model_config.py +++ b/torax/_src/torax_pydantic/model_config.py @@ -16,7 +16,7 @@ import copy import logging -from typing import Any, Mapping +from typing import Any, Mapping, Self import numpy as np import pydantic @@ -44,8 +44,6 @@ from torax._src.torax_pydantic import file_restart as file_restart_pydantic_model from torax._src.torax_pydantic import torax_pydantic from torax._src.transport_model import pydantic_model as transport_model_pydantic_model -import typing_extensions -from typing_extensions import Self class ToraxConfig(torax_pydantic.BaseModelFrozen): @@ -141,7 +139,7 @@ def _defaults(cls, data: dict[str, Any]) -> dict[str, Any]: return configurable_data @pydantic.model_validator(mode='after') - def _check_fields(self) -> typing_extensions.Self: + def _check_fields(self) -> Self: core_transport_models = self.transport.core_transport_models.values() pedestal_transport_models = ( self.transport.pedestal_transport_models.values() @@ -155,12 +153,9 @@ def _check_fields(self) -> typing_extensions.Self: self.solver, solver_pydantic_model.LinearThetaMethod ) - # pylint: disable=g-long-ternary - # pylint: disable=attribute-error initial_guess_mode_is_linear = ( - False - if using_linear_solver - else self.solver.initial_guess_mode == enums.InitialGuessMode.LINEAR # pyrefly: ignore[missing-attribute] + getattr(self.solver, 'initial_guess_mode', None) + == enums.InitialGuessMode.LINEAR ) if ( @@ -182,7 +177,7 @@ def _check_fields(self) -> typing_extensions.Self: return self @pydantic.model_validator(mode='after') - def _check_psidot_and_evolve_current(self) -> typing_extensions.Self: + def _check_psidot_and_evolve_current(self) -> Self: """Warns if psidot is provided but evolve_current is True.""" if ( self.profile_conditions.psidot is not None @@ -197,7 +192,7 @@ def _check_psidot_and_evolve_current(self) -> typing_extensions.Self: return self @pydantic.model_validator(mode='after') - def _check_pedestal_with_non_uniform_grid(self) -> typing_extensions.Self: + def _check_pedestal_with_non_uniform_grid(self) -> Self: """Warns if a pedestal and non-uniform grid are used.""" if self.pedestal.model_name != 'no_pedestal': face_centers = self.geometry.get_face_centers() @@ -217,7 +212,7 @@ def _check_pedestal_with_non_uniform_grid(self) -> typing_extensions.Self: @pydantic.model_validator(mode='after') def _validate_pedestal_mode_and_internal_boundary_conditions( self, - ) -> typing_extensions.Self: + ) -> Self: """Validates that internal boundary conditions are not used with ADAPTIVE_TRANSPORT.""" ibc = self.profile_conditions.internal_boundary_conditions if ( @@ -233,7 +228,7 @@ def _validate_pedestal_mode_and_internal_boundary_conditions( return self @pydantic.model_validator(mode='after') - def _check_edge_with_circular_geometry(self) -> typing_extensions.Self: + def _check_edge_with_circular_geometry(self) -> Self: """Validates that edge models are not used with CircularGeometry.""" if ( self.edge is not None @@ -247,7 +242,7 @@ def _check_edge_with_circular_geometry(self) -> typing_extensions.Self: @pydantic.model_validator(mode='after') def _validate_extended_lengyel_and_impurity_mode( self, - ) -> typing_extensions.Self: + ) -> Self: """Ensures Extended Lengyel uses n_e_ratios impurity mode.""" if ( isinstance( @@ -264,7 +259,7 @@ def _validate_extended_lengyel_and_impurity_mode( return self @pydantic.model_validator(mode='after') - def _validate_edge_diverted_status(self) -> typing_extensions.Self: + def _validate_edge_diverted_status(self) -> Self: """Validates diverted status configuration in edge model. Ensures that `diverted` is handled correctly based on geometry type: @@ -296,7 +291,7 @@ def _validate_edge_diverted_status(self) -> typing_extensions.Self: return self @pydantic.model_validator(mode='after') - def _validate_edge_core_impurity_consistency(self) -> typing_extensions.Self: + def _validate_edge_core_impurity_consistency(self) -> Self: """Validates consistency between plasma composition and edge impurities.""" if isinstance( self.edge, @@ -338,7 +333,7 @@ def _validate_edge_core_impurity_consistency(self) -> typing_extensions.Self: return self @pydantic.model_validator(mode='after') - def _validate_nonzero_n_e_ratios_at_lcfs(self) -> typing_extensions.Self: + def _validate_nonzero_n_e_ratios_at_lcfs(self) -> Self: """Validates that n_e_ratio profiles are non-zero at the LCFS. When the extended Lengyel edge model is active, core impurity profiles @@ -475,7 +470,7 @@ def torax_version(self) -> str: return version.TORAX_VERSION @pydantic.model_validator(mode='after') - def _validate_toric_nn_he3_presence(self) -> typing_extensions.Self: + def _validate_toric_nn_he3_presence(self) -> Self: """Validates that He3 is present in plasma composition if ToricNN is used. The ToricNN model currently only supports He3 minority heating, so He3 must diff --git a/torax/_src/torax_pydantic/tests/interpolated_param_1d_test.py b/torax/_src/torax_pydantic/tests/interpolated_param_1d_test.py index 288792eaa..3b7ceb618 100644 --- a/torax/_src/torax_pydantic/tests/interpolated_param_1d_test.py +++ b/torax/_src/torax_pydantic/tests/interpolated_param_1d_test.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +from typing import Annotated from absl.testing import absltest from absl.testing import parameterized import chex @@ -23,7 +24,6 @@ from torax._src.geometry import circular_geometry from torax._src.torax_pydantic import interpolated_param_1d from torax._src.torax_pydantic import torax_pydantic -import typing_extensions import xarray as xr RHO_NORM = 'rho_norm' @@ -212,7 +212,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): @parameterized.named_parameters( dict( testcase_name='gt_valid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalar, torax_pydantic.scalar_bounds_validator(gt=1.0), ], @@ -221,7 +221,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): ), dict( testcase_name='gt_equal_invalid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalar, torax_pydantic.scalar_bounds_validator(gt=1.0), ], @@ -231,7 +231,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): ), dict( testcase_name='ge_equal_valid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalar, torax_pydantic.scalar_bounds_validator(ge=1.0), ], @@ -240,7 +240,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): ), dict( testcase_name='lt_valid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalar, torax_pydantic.scalar_bounds_validator(lt=5.0), ], @@ -249,7 +249,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): ), dict( testcase_name='lt_equal_invalid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalar, torax_pydantic.scalar_bounds_validator(lt=5.0), ], @@ -259,7 +259,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): ), dict( testcase_name='le_equal_valid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalar, torax_pydantic.scalar_bounds_validator(le=5.0), ], @@ -268,7 +268,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): ), dict( testcase_name='interval_valid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalar, torax_pydantic.scalar_bounds_validator( gt=1.0, lt=10.0 @@ -279,7 +279,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): ), dict( testcase_name='interval_below_invalid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalar, torax_pydantic.scalar_bounds_validator( gt=1.0, lt=10.0 @@ -291,7 +291,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): ), dict( testcase_name='interval_above_invalid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalar, torax_pydantic.scalar_bounds_validator( gt=1.0, lt=10.0 @@ -303,7 +303,7 @@ class TestModel(torax_pydantic.BaseModelFrozen): ), dict( testcase_name='step_mode_valid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ torax_pydantic.TimeVaryingScalarStep, torax_pydantic.scalar_bounds_validator(ge=0.0), ], diff --git a/torax/_src/torax_pydantic/tests/interpolated_param_2d_test.py b/torax/_src/torax_pydantic/tests/interpolated_param_2d_test.py index 7bf9209aa..9df1956fe 100644 --- a/torax/_src/torax_pydantic/tests/interpolated_param_2d_test.py +++ b/torax/_src/torax_pydantic/tests/interpolated_param_2d_test.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +from typing import Annotated from absl.testing import absltest from absl.testing import parameterized import chex @@ -23,7 +24,6 @@ from torax._src.geometry import circular_geometry from torax._src.torax_pydantic import interpolated_param_2d from torax._src.torax_pydantic import model_base -import typing_extensions import xarray as xr RHO_NORM = 'rho_norm' @@ -316,7 +316,7 @@ class TestModel(model_base.BaseModelFrozen): @parameterized.named_parameters( dict( testcase_name='gt_valid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ interpolated_param_2d.TimeVaryingArray, interpolated_param_2d.array_bounds_validator(gt=1.0), ], @@ -325,7 +325,7 @@ class TestModel(model_base.BaseModelFrozen): ), dict( testcase_name='gt_equal_invalid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ interpolated_param_2d.TimeVaryingArray, interpolated_param_2d.array_bounds_validator(gt=1.0), ], @@ -335,7 +335,7 @@ class TestModel(model_base.BaseModelFrozen): ), dict( testcase_name='ge_equal_valid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ interpolated_param_2d.TimeVaryingArray, interpolated_param_2d.array_bounds_validator(ge=1.0), ], @@ -344,7 +344,7 @@ class TestModel(model_base.BaseModelFrozen): ), dict( testcase_name='interval_valid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ interpolated_param_2d.TimeVaryingArray, interpolated_param_2d.array_bounds_validator( gt=0.0, lt=5.0 @@ -355,7 +355,7 @@ class TestModel(model_base.BaseModelFrozen): ), dict( testcase_name='interval_above_invalid', - field_type=typing_extensions.Annotated[ + field_type=Annotated[ interpolated_param_2d.TimeVaryingArray, interpolated_param_2d.array_bounds_validator( gt=0.0, lt=5.0 diff --git a/torax/_src/torax_pydantic/torax_pydantic.py b/torax/_src/torax_pydantic/torax_pydantic.py index 0d9b7fc6f..fa31c6406 100644 --- a/torax/_src/torax_pydantic/torax_pydantic.py +++ b/torax/_src/torax_pydantic/torax_pydantic.py @@ -15,14 +15,13 @@ """Pydantic utilities and base classes.""" import functools -from typing import TypeAlias +from typing import Annotated, TypeAlias import pydantic from torax._src.torax_pydantic import interpolated_param_1d from torax._src.torax_pydantic import interpolated_param_2d from torax._src.torax_pydantic import model_base from torax._src.torax_pydantic import pydantic_types -from typing_extensions import Annotated TIME_INVARIANT = model_base.TIME_INVARIANT JAX_STATIC = model_base.JAX_STATIC diff --git a/torax/_src/transport_model/component.py b/torax/_src/transport_model/component.py index 4cdf66bef..7436bbcd1 100644 --- a/torax/_src/transport_model/component.py +++ b/torax/_src/transport_model/component.py @@ -21,7 +21,7 @@ import abc import dataclasses -from typing import ClassVar, Mapping, Sequence +from typing import ClassVar, Mapping, Sequence, cast import immutabledict import jax @@ -66,23 +66,23 @@ class TurbulentTransport: v_face_el_tem: (Optional) TEM contribution for electron convection. """ - chi_face_ion: jax.Array - chi_face_el: jax.Array - d_face_el: jax.Array - v_face_el: jax.Array - chi_face_el_bohm: jax.Array | None = None - chi_face_el_gyrobohm: jax.Array | None = None - chi_face_ion_bohm: jax.Array | None = None - chi_face_ion_gyrobohm: jax.Array | None = None - chi_face_ion_itg: jax.Array | None = None - chi_face_ion_tem: jax.Array | None = None - chi_face_el_itg: jax.Array | None = None - chi_face_el_tem: jax.Array | None = None - chi_face_el_etg: jax.Array | None = None - d_face_el_itg: jax.Array | None = None - d_face_el_tem: jax.Array | None = None - v_face_el_itg: jax.Array | None = None - v_face_el_tem: jax.Array | None = None + chi_face_ion: array_typing.FloatVectorFace + chi_face_el: array_typing.FloatVectorFace + d_face_el: array_typing.FloatVectorFace + v_face_el: array_typing.FloatVectorFace + chi_face_el_bohm: array_typing.FloatVectorFace | None = None + chi_face_el_gyrobohm: array_typing.FloatVectorFace | None = None + chi_face_ion_bohm: array_typing.FloatVectorFace | None = None + chi_face_ion_gyrobohm: array_typing.FloatVectorFace | None = None + chi_face_ion_itg: array_typing.FloatVectorFace | None = None + chi_face_ion_tem: array_typing.FloatVectorFace | None = None + chi_face_el_itg: array_typing.FloatVectorFace | None = None + chi_face_el_tem: array_typing.FloatVectorFace | None = None + chi_face_el_etg: array_typing.FloatVectorFace | None = None + d_face_el_itg: array_typing.FloatVectorFace | None = None + d_face_el_tem: array_typing.FloatVectorFace | None = None + v_face_el_itg: array_typing.FloatVectorFace | None = None + v_face_el_tem: array_typing.FloatVectorFace | None = None @dataclasses.dataclass(frozen=True, eq=False) @@ -181,7 +181,9 @@ def zero_out_disabled_channels( to_replace = {} for channel_name, config in self.CHANNEL_CONFIG.items(): - disable_flag = getattr(transport_runtime_params, config['disable_flag']) # pyrefly: ignore[bad-argument-type] + disable_flag = getattr( + transport_runtime_params, cast(str, config['disable_flag']) + ) # Handle main channel val = getattr(transport_coeffs, channel_name) diff --git a/torax/_src/transport_model/pydantic_model.py b/torax/_src/transport_model/pydantic_model.py index 0eb235f3a..0e9da2629 100644 --- a/torax/_src/transport_model/pydantic_model.py +++ b/torax/_src/transport_model/pydantic_model.py @@ -16,7 +16,7 @@ import copy import dataclasses -from typing import Annotated, Any, Literal, Mapping, Sequence +from typing import Annotated, Any, Literal, Mapping, Self, Sequence from absl import logging import chex from fusion_surrogates.qlknn.models import registry @@ -36,7 +36,6 @@ from torax._src.transport_model import tglfnn_ukaea_transport_model from torax._src.transport_model import transport_model from torax._src.transport_model.tglf import tglf_transport_model -import typing_extensions def _resolve_qlknn_model_name(model_name: str, model_path: str) -> str: @@ -493,12 +492,12 @@ class TransportModel(torax_pydantic.BaseModelFrozen): str, ComponentTransportModelConfig ] = pydantic.Field( default_factory=dict - ) # pyrefly: ignore[invalid-annotation] + ) pedestal_transport_models: dict[ str, ComponentTransportModelConfig ] = pydantic.Field( default_factory=dict - ) # pyrefly: ignore[invalid-annotation] + ) smoothing_zones: Sequence[SmoothingZone] = pydantic.Field( default_factory=list ) @@ -553,7 +552,7 @@ def build_runtime_params( ) @pydantic.model_validator(mode='after') - def _check_smoothing_width_minimum(self) -> typing_extensions.Self: + def _check_smoothing_width_minimum(self) -> Self: smoothing_widths = [z.smoothing_width for z in self.smoothing_zones] + [ self.smoothing_width ] @@ -572,7 +571,7 @@ def _check_smoothing_width_minimum(self) -> typing_extensions.Self: return self @pydantic.model_validator(mode='after') - def _check_fields(self) -> typing_extensions.Self: + def _check_fields(self) -> Self: if not self.chi_min < self.chi_max: raise ValueError('chi_min must be less than chi_max.') if not self.D_e_min < self.D_e_max: @@ -591,12 +590,12 @@ def _check_fields(self) -> typing_extensions.Self: return self @pydantic.model_validator(mode='after') - def _check_unique_overwrites_core(self) -> typing_extensions.Self: + def _check_unique_overwrites_core(self) -> Self: _validate_unique_overwrites(self.core_transport_models, 'core') return self @pydantic.model_validator(mode='after') - def _check_unique_overwrites_pedestal(self) -> typing_extensions.Self: + def _check_unique_overwrites_pedestal(self) -> Self: _validate_unique_overwrites(self.pedestal_transport_models, 'pedestal') return self diff --git a/torax/_src/transport_model/pydantic_model_base.py b/torax/_src/transport_model/pydantic_model_base.py index ee2aaf064..6e9e373d2 100644 --- a/torax/_src/transport_model/pydantic_model_base.py +++ b/torax/_src/transport_model/pydantic_model_base.py @@ -15,7 +15,7 @@ """Base pydantic config for Transport models.""" import abc -from typing import Annotated +from typing import Annotated, Self import chex import numpy as np @@ -25,7 +25,6 @@ from torax._src.transport_model import component from torax._src.transport_model import enums from torax._src.transport_model import runtime_params -import typing_extensions # pylint: disable=invalid-name @@ -80,7 +79,7 @@ class ComponentTransportBase(torax_pydantic.BaseModelFrozen, abc.ABC): ) @pydantic.model_validator(mode='after') - def _check_fields(self) -> typing_extensions.Self: + def _check_fields(self) -> Self: # For the time-varying parameter pair (rho_min, rho_max), we have relative # magnitude constraints that must hold at all times. We validate this by # checking the inequality at the combined time points (knots) of the pair. diff --git a/torax/_src/transport_model/qlknn_10d.py b/torax/_src/transport_model/qlknn_10d.py index 1c08e51ac..ebce42521 100644 --- a/torax/_src/transport_model/qlknn_10d.py +++ b/torax/_src/transport_model/qlknn_10d.py @@ -16,7 +16,7 @@ from collections.abc import Mapping import json import os -from typing import Any, Callable, Final +from typing import Any, Callable, Final, Self import flax.linen as nn import immutabledict @@ -26,7 +26,6 @@ from torax._src import jax_utils from torax._src.transport_model import base_qlknn_model from torax._src.transport_model import qualikiz_based_transport_model -import typing_extensions # Internal import. # Internal import. @@ -118,7 +117,7 @@ def __call__( return outputs @classmethod - def from_json(cls, json_file) -> typing_extensions.Self: + def from_json(cls, json_file) -> Self: with open(json_file) as file_: model_dict = json.load(file_) return cls(model_dict) diff --git a/torax/_src/transport_model/qualikiz_based_transport_model.py b/torax/_src/transport_model/qualikiz_based_transport_model.py index e5ca3f565..a414a0d10 100644 --- a/torax/_src/transport_model/qualikiz_based_transport_model.py +++ b/torax/_src/transport_model/qualikiz_based_transport_model.py @@ -123,7 +123,7 @@ def _prepare_qualikiz_inputs( # gyrobohm diffusivity # (defined here with Lref=a_minor due to QLKNN training set normalization) chiGB = quasilinear_transport_model.calculate_chiGB( - reference_temperature=core_profiles.T_i.face_value(), # pyrefly: ignore[bad-argument-type] + reference_temperature=core_profiles.T_i.face_value(), reference_magnetic_field=geo.B_0, reference_mass=core_profiles.A_i, reference_length=geo.a_minor, @@ -306,9 +306,9 @@ def _calc_gamma_E_SI(v_ExB_component): q=q, smag=smag, x=x, - Ti_Te=Ti_Te, # pyrefly: ignore[bad-argument-type] + Ti_Te=Ti_Te, log_nu_star_face=log_nu_star_face, - normni=normni, # pyrefly: ignore[bad-argument-type] + normni=normni, chiGB=chiGB, Rmaj=geo.R_major, Rmin=geo.a_minor, diff --git a/torax/_src/transport_model/quasilinear_transport_model.py b/torax/_src/transport_model/quasilinear_transport_model.py index ea47b9403..2563ec199 100644 --- a/torax/_src/transport_model/quasilinear_transport_model.py +++ b/torax/_src/transport_model/quasilinear_transport_model.py @@ -16,6 +16,7 @@ from collections.abc import Mapping import dataclasses import functools +from typing import Self import chex from fusion_surrogates.fast_ion_stabilization import fast_ion_model from fusion_surrogates.fast_ion_stabilization.models import registry as fi_registry @@ -29,7 +30,6 @@ from torax._src.geometry import geometry from torax._src.transport_model import component from torax._src.transport_model import runtime_params as runtime_params_lib -import typing_extensions @jax.tree_util.register_dataclass @@ -58,7 +58,7 @@ def from_profiles( radial_face_coordinate: jnp.ndarray, reference_length: jnp.ndarray, two_point_mask: array_typing.BoolVectorFace | None = None, - ) -> typing_extensions.Self: + ) -> Self: """Calculates the normalized logarithmic gradients.""" gradients = {} for name, profile in { @@ -313,8 +313,8 @@ def _load_fi_stabilization_model(model: str): def _compute_fast_ion_stabilization_factor( core_profiles: state.CoreProfiles, - smag: jax.Array, - q: jax.Array, + smag: array_typing.Array, + q: array_typing.Array, normalized_logarithmic_gradients: NormalizedLogarithmicGradients, model_map: dict[str, str] | None = None, ) -> jax.Array: @@ -367,8 +367,8 @@ def _compute_fast_ion_stabilization_factor( def apply_fast_ion_stabilization( core_profiles: state.CoreProfiles, - smag: jax.Array, - q: jax.Array, + smag: array_typing.Array, + q: array_typing.Array, normalized_logarithmic_gradients: NormalizedLogarithmicGradients, transport: RuntimeParams, ) -> jax.Array: @@ -457,7 +457,9 @@ def _make_core_transport( # Effective D / Effective V approach. # For small density gradients or up-gradient transport, set pure effective # convection. Otherwise pure effective diffusion. - def DV_effective_approach() -> tuple[jax.Array, jax.Array]: + def DV_effective_approach() -> ( + tuple[array_typing.FloatVectorFace, array_typing.FloatVectorFace] + ): # The geo.rho_b is to unnormalize the face_grad. Deff = -pfe_SI / ( core_profiles.n_e.face_grad(two_point_mask=two_point_mask) @@ -482,7 +484,9 @@ def DV_effective_approach() -> tuple[jax.Array, jax.Array]: # Scaled D approach. Scale electron diffusivity to electron heat # conductivity (this has some physical motivations), # and set convection to then match total particle transport - def Dscaled_approach() -> tuple[jax.Array, jax.Array]: + def Dscaled_approach() -> ( + tuple[array_typing.FloatVectorFace, array_typing.FloatVectorFace] + ): chex.assert_rank(pfe, 1) d_face_el = chi_face_el v_face_el = ( @@ -493,7 +497,7 @@ def Dscaled_approach() -> tuple[jax.Array, jax.Array]: * geo.g1_over_vpr2_face * geo.rho_b**2 ) / (geo.g0_over_vpr_face * geo.rho_b) - return d_face_el, v_face_el # pyrefly: ignore[bad-return] + return d_face_el, v_face_el d_face_el, v_face_el = jax.lax.cond( transport.DV_effective, @@ -501,8 +505,8 @@ def Dscaled_approach() -> tuple[jax.Array, jax.Array]: Dscaled_approach, ) return component.TurbulentTransport( - chi_face_ion=chi_face_ion, # pyrefly: ignore[bad-argument-type] - chi_face_el=chi_face_el, # pyrefly: ignore[bad-argument-type] + chi_face_ion=chi_face_ion, + chi_face_el=chi_face_el, d_face_el=d_face_el, v_face_el=v_face_el, ) diff --git a/torax/_src/transport_model/tests/qualikiz_based_transport_model_test.py b/torax/_src/transport_model/tests/qualikiz_based_transport_model_test.py index 959f3a310..106b7c864 100644 --- a/torax/_src/transport_model/tests/qualikiz_based_transport_model_test.py +++ b/torax/_src/transport_model/tests/qualikiz_based_transport_model_test.py @@ -35,6 +35,7 @@ from torax._src.transport_model import pydantic_model_base as transport_pydantic_model_base from torax._src.transport_model import qualikiz_based_transport_model from torax._src.transport_model import register_model +from torax._src.transport_model import runtime_params as transport_runtime_params_lib def setUpModule(): @@ -271,9 +272,11 @@ def prepare_qualikiz_inputs( # pylint: enable=invalid-name - def call_implementation( # pyrefly: ignore[bad-override] + def call_implementation( self, - transport_runtime_params: qualikiz_based_transport_model.RuntimeParams, + transport_runtime_params: ( + transport_runtime_params_lib.ComponentRuntimeParams + ), runtime_params: runtime_params_lib.RuntimeParams, geo: geometry.Geometry, core_profiles: state.CoreProfiles, diff --git a/torax/_src/transport_model/tests/quasilinear_transport_model_test.py b/torax/_src/transport_model/tests/quasilinear_transport_model_test.py index ca8b125e8..a2b6e6b37 100644 --- a/torax/_src/transport_model/tests/quasilinear_transport_model_test.py +++ b/torax/_src/transport_model/tests/quasilinear_transport_model_test.py @@ -518,14 +518,20 @@ class FakeQuasilinearTransportModel( ): """Fake QuasilinearTransportModel for testing purposes.""" - def call_implementation( # pyrefly: ignore[bad-override] + def call_implementation( self, - transport_runtime_params: quasilinear_transport_model.RuntimeParams, + transport_runtime_params: ( + transport_model_runtime_params.ComponentRuntimeParams + ), runtime_params: runtime_params_lib.RuntimeParams, geo: geometry.Geometry, core_profiles: state.CoreProfiles, two_point_mask: array_typing.BoolVectorFace, ) -> component.TurbulentTransport: + assert isinstance( + transport_runtime_params, + quasilinear_transport_model.RuntimeParams, + ) quasilinear_inputs = quasilinear_transport_model.QuasilinearInputs( chiGB=np.array(4.0), Rmin=np.array(0.5), diff --git a/torax/_src/transport_model/tests/tglf_based_transport_model_test.py b/torax/_src/transport_model/tests/tglf_based_transport_model_test.py index ed9a5c580..73d2a84cf 100644 --- a/torax/_src/transport_model/tests/tglf_based_transport_model_test.py +++ b/torax/_src/transport_model/tests/tglf_based_transport_model_test.py @@ -34,6 +34,7 @@ from torax._src.transport_model import component from torax._src.transport_model import pydantic_model_base as transport_pydantic_model_base from torax._src.transport_model import register_model +from torax._src.transport_model import runtime_params as transport_runtime_params_lib from torax._src.transport_model import tglf_based_transport_model from torax._src.transport_model.tglf import tglf2py @@ -245,9 +246,11 @@ def prepare_tglf_inputs( # pylint: enable=invalid-name - def call_implementation( # pyrefly: ignore[bad-override] + def call_implementation( self, - transport_runtime_params: tglf_based_transport_model.RuntimeParams, + transport_runtime_params: ( + transport_runtime_params_lib.ComponentRuntimeParams + ), runtime_params: runtime_params_lib.RuntimeParams, geo: geometry.Geometry, core_profiles: state.CoreProfiles, diff --git a/torax/_src/transport_model/tglf_based_transport_model.py b/torax/_src/transport_model/tglf_based_transport_model.py index 7f0c87ed3..588d7552e 100644 --- a/torax/_src/transport_model/tglf_based_transport_model.py +++ b/torax/_src/transport_model/tglf_based_transport_model.py @@ -28,7 +28,6 @@ from torax._src.physics import rotation from torax._src.transport_model import component from torax._src.transport_model import quasilinear_transport_model -from typing_extensions import override @jax.tree_util.register_dataclass @@ -131,7 +130,7 @@ class TGLFInputs(quasilinear_transport_model.QuasilinearInputs): class TGLFBasedTransportModel( - quasilinear_transport_model.QuasilinearTransportModel + component.ComponentTransportModel ): """Base class for TGLF-based transport models.""" @@ -231,7 +230,7 @@ def _prepare_tglf_inputs( # avoid being swamped by the eps in the denominator. rho_s = ( math_utils.safe_divide( - num=m_D * c_s, # pyrefly: ignore[bad-argument-type] + num=m_D * c_s, denom=B_unit, eps=1e-7, ) @@ -437,7 +436,7 @@ def _get_v_ExB_shear( lref_over_lti = quasilinear_transport_model.apply_fast_ion_stabilization( core_profiles=core_profiles, smag=smag, - q=core_profiles.q_face, # pyrefly: ignore[bad-argument-type] + q=core_profiles.q_face, normalized_logarithmic_gradients=normalized_log_gradients, transport=transport, ) @@ -470,12 +469,12 @@ def _get_v_ExB_shear( AS_1=n_e_over_n_e, ZS_2=core_profiles.Z_i_face, MASS_2=m_i_over_m_D, - TAUS_2=T_i_over_T_e, # pyrefly: ignore[bad-argument-type] - AS_2=n_i_over_n_e, # pyrefly: ignore[bad-argument-type] + TAUS_2=T_i_over_T_e, + AS_2=n_i_over_n_e, ZS_3=core_profiles.Z_impurity_face, MASS_3=m_imp_over_m_D, - TAUS_3=T_imp_over_T_e, # pyrefly: ignore[bad-argument-type] - AS_3=n_impurity_over_n_e, # pyrefly: ignore[bad-argument-type] + TAUS_3=T_imp_over_T_e, + AS_3=n_impurity_over_n_e, RLNS_1=normalized_log_gradients.lref_over_lne, RLNS_2=normalized_log_gradients.lref_over_lni0, RLNS_3=normalized_log_gradients.lref_over_lni1, @@ -486,23 +485,22 @@ def _get_v_ExB_shear( RMAJ_LOC=r_major / a, DRMAJDX_LOC=dr_major, # pyrefly: ignore[bad-argument-type] Q_LOC=core_profiles.q_face, - Q_PRIME_LOC=q_prime, # pyrefly: ignore[bad-argument-type] + Q_PRIME_LOC=q_prime, XNUE=normalized_nu_ee, - DEBYE=normalized_debye, # pyrefly: ignore[bad-argument-type] + DEBYE=normalized_debye, KAPPA_LOC=kappa, S_KAPPA_LOC=kappa_shear, # pyrefly: ignore[bad-argument-type] DELTA_LOC=geo.delta_face, S_DELTA_LOC=delta_shear, # pyrefly: ignore[bad-argument-type] - BETAE=beta_e, # pyrefly: ignore[bad-argument-type] - P_PRIME_LOC=p_prime, # pyrefly: ignore[bad-argument-type] + BETAE=beta_e, + P_PRIME_LOC=p_prime, ZEFF=core_profiles.Z_eff_face, - Q_GB=Q_GB, # pyrefly: ignore[bad-argument-type] - GAMMA_GB=Gamma_GB, # pyrefly: ignore[bad-argument-type] + Q_GB=Q_GB, + GAMMA_GB=Gamma_GB, VEXB_SHEAR=v_ExB_shear, ) - @override - def _make_core_transport( # pyrefly: ignore[bad-override] + def _make_core_transport( self, electron_heat_flux_GB: jax.Array, ion_heat_flux_GB: jax.Array, @@ -581,8 +579,8 @@ def _make_core_transport( # pyrefly: ignore[bad-override] v_face_el = jnp.where(V_eff_mask, V_eff, 0.0) return component.TurbulentTransport( - chi_face_ion=chi_i, # pyrefly: ignore[bad-argument-type] - chi_face_el=chi_e, # pyrefly: ignore[bad-argument-type] + chi_face_ion=chi_i, + chi_face_el=chi_e, d_face_el=d_face_el, v_face_el=v_face_el, ) diff --git a/torax/_src/transport_model/transport_coefficients_builder.py b/torax/_src/transport_model/transport_coefficients_builder.py index 5b4696ca9..a261470a9 100644 --- a/torax/_src/transport_model/transport_coefficients_builder.py +++ b/torax/_src/transport_model/transport_coefficients_builder.py @@ -144,7 +144,7 @@ def calculate_all_transport_coeffs( core_transport = state.CoreTransport( **dataclasses.asdict(turbulent_transport_coeffs), - **dataclasses.asdict(neoclassical_transport_coeffs), # pyrefly: ignore[bad-argument-type] + **dataclasses.asdict(neoclassical_transport_coeffs), **dataclasses.asdict(pereverzev_transport_coeffs), ) diff --git a/torax/_src/transport_model/transport_coeffs.py b/torax/_src/transport_model/transport_coeffs.py index aa45bc7e0..b2bd8cc9d 100644 --- a/torax/_src/transport_model/transport_coeffs.py +++ b/torax/_src/transport_model/transport_coeffs.py @@ -15,6 +15,7 @@ """Transport coefficient data structures.""" import dataclasses +from typing import Self import jax from jax import numpy as jnp @@ -22,7 +23,6 @@ from torax._src.geometry import geometry from torax._src.output_tools import output_grid_context from torax._src.output_tools import output_keys -import typing_extensions # pylint: disable=invalid-name @@ -44,7 +44,7 @@ class TransportCoeffs: v_face_el: array_typing.FloatVectorFace @classmethod - def zeros(cls, geo: geometry.Geometry) -> typing_extensions.Self: + def zeros(cls, geo: geometry.Geometry) -> Self: """Returns a TransportCoeffs with all zeros.""" zeros = jnp.zeros_like(geo.rho_face_norm) return cls( @@ -54,7 +54,7 @@ def zeros(cls, geo: geometry.Geometry) -> typing_extensions.Self: v_face_el=zeros, ) - def __add__(self, other: typing_extensions.Self) -> typing_extensions.Self: + def __add__(self, other: Self) -> Self: """Adds two TransportCoeffs channel-by-channel.""" return self.__class__( chi_face_ion=self.chi_face_ion + other.chi_face_ion, diff --git a/torax/_src/tridiagonal.py b/torax/_src/tridiagonal.py index feb6119d5..6a02b3938 100644 --- a/torax/_src/tridiagonal.py +++ b/torax/_src/tridiagonal.py @@ -14,6 +14,7 @@ """Tridiagonal matrix representations and operations.""" +from collections.abc import Iterable import dataclasses import enum @@ -23,7 +24,6 @@ import jaxtyping as jt from torax._src import array_typing from torax._src import jax_utils -import typing_extensions @enum.unique @@ -54,8 +54,8 @@ def to_dense(self) -> jt.Float[array_typing.Array, 'size size']: + jnp.diag(self.below, -1) ) - def __add__(self, other: typing_extensions.Self) -> typing_extensions.Self: - return TriDiagonal( # pyrefly: ignore[bad-return] + def __add__(self, other: 'TriDiagonal') -> 'TriDiagonal': + return TriDiagonal( diagonal=self.diagonal + other.diagonal, above=self.above + other.above, below=self.below + other.below, @@ -97,8 +97,8 @@ def block_size(self) -> int: """Size of each block.""" return self.diagonal.shape[1] - def __add__(self, other: typing_extensions.Self) -> typing_extensions.Self: - return BlockTriDiagonal( # pyrefly: ignore[bad-return] + def __add__(self, other: 'BlockTriDiagonal') -> 'BlockTriDiagonal': + return BlockTriDiagonal( lower=self.lower + other.lower, diagonal=self.diagonal + other.diagonal, upper=self.upper + other.upper, @@ -146,7 +146,7 @@ def from_diagonal( @classmethod def from_tridiagonals( cls, - tridiagonals: typing_extensions.Iterable[TriDiagonal], + tridiagonals: Iterable[TriDiagonal], ) -> 'BlockTriDiagonal': """Creates a BlockTriDiagonal from an iterable of per-channel TriDiagonals. @@ -162,7 +162,7 @@ def from_tridiagonals( tridiagonals_seq = tuple(tridiagonals) stacked = jax.tree.map( lambda *args: jnp.stack(args, axis=1), *tridiagonals_seq - ) + ) return cls( lower=stacked.below[..., None, :] * jnp.eye(stacked.below.shape[-1], dtype=stacked.below.dtype), diff --git a/torax/_src/version.py b/torax/_src/version.py index a7535583c..c81bc2b53 100644 --- a/torax/_src/version.py +++ b/torax/_src/version.py @@ -21,7 +21,8 @@ def _version_as_tuple(version_str: str) -> tuple[int, int, int]: - return tuple(int(i) for i in version_str.split(".") if i.isdigit()) # pyrefly: ignore[bad-return] + major, minor, patch = (int(i) for i in version_str.split(".") if i.isdigit()) + return (major, minor, patch) TORAX_VERSION_INFO: Final[tuple[int, int, int]] = _version_as_tuple( diff --git a/torax/experimental/gas_puff_feedback_source.py b/torax/experimental/gas_puff_feedback_source.py index 1acbd9927..9370d29b0 100644 --- a/torax/experimental/gas_puff_feedback_source.py +++ b/torax/experimental/gas_puff_feedback_source.py @@ -83,9 +83,9 @@ def calc_puff_feedback_source( match source_params.average_type: case AverageType.LINE: - current_avg_n_e = math_utils.line_average(core_profiles.n_e.value, geo) # pyrefly: ignore[bad-argument-type] + current_avg_n_e = math_utils.line_average(core_profiles.n_e.value, geo) case AverageType.VOLUME: - current_avg_n_e = math_utils.volume_average(core_profiles.n_e.value, geo) # pyrefly: ignore[bad-argument-type] + current_avg_n_e = math_utils.volume_average(core_profiles.n_e.value, geo) case _ as unknown: raise ValueError(f'Unknown average type: {unknown}') @@ -98,8 +98,8 @@ def calc_puff_feedback_source( return ( formulas.exponential_profile( decay_start=1.0, - width=source_params.puff_decay_length, # pyrefly: ignore[bad-argument-type] - total=S_total, # pyrefly: ignore[bad-argument-type] + width=source_params.puff_decay_length, + total=S_total, geo=geo, ), ) @@ -161,7 +161,7 @@ class GasPuffFeedbackSourceConfig(base.SourceModelBase): @property def model_func(self) -> source.SourceProfileFunction: - return calc_puff_feedback_source # pyrefly: ignore[bad-return] + return calc_puff_feedback_source def build_runtime_params( self, diff --git a/torax/experimental/tests/gas_puff_feedback_source_test.py b/torax/experimental/tests/gas_puff_feedback_source_test.py index 57f9a6cc9..d99e0d970 100644 --- a/torax/experimental/tests/gas_puff_feedback_source_test.py +++ b/torax/experimental/tests/gas_puff_feedback_source_test.py @@ -64,7 +64,7 @@ def test_feedback_mode(self): neoclassical_models=neoclassical_models, ) - initial_line_avg = math_utils.line_average(core_profiles.n_e.value, geo) # pyrefly: ignore[bad-argument-type] + initial_line_avg = math_utils.line_average(core_profiles.n_e.value, geo) # Rebuild with specific requested value config['sources']['gas_puff']['model_name'] = 'feedback' diff --git a/torax/run_simulation_main.py b/torax/run_simulation_main.py index e20905b20..7fcd5cd8f 100644 --- a/torax/run_simulation_main.py +++ b/torax/run_simulation_main.py @@ -229,20 +229,16 @@ def _change_config( return torax_config, config_path -def _get_yes_or_no() -> bool: # pyrefly: ignore[bad-return] +def _get_yes_or_no() -> bool: """Returns a boolean indicating yes depending on user input.""" - input_text = None - while input_text is None: - input_text = input(Y_N_PROMPT) - input_text = input_text.lower().strip() - if input_text not in ('y', 'n'): - simulation_app.log_to_stdout( - 'Unrecognized input. Try again.', - color=simulation_app.AnsiColors.YELLOW, - ) - input_text = None - else: + while True: + input_text = input(Y_N_PROMPT).lower().strip() + if input_text in ('y', 'n'): return input_text == 'y' + simulation_app.log_to_stdout( + 'Unrecognized input. Try again.', + color=simulation_app.AnsiColors.YELLOW, + ) def _toggle_log_progress(log_sim_progress: bool) -> bool: diff --git a/torax/tests/sim_time_dependence_test.py b/torax/tests/sim_time_dependence_test.py index 46b89157e..f363e32a9 100644 --- a/torax/tests/sim_time_dependence_test.py +++ b/torax/tests/sim_time_dependence_test.py @@ -15,6 +15,7 @@ """Tests torax.sim for handling time dependent input runtime params.""" import dataclasses +import functools from typing import Annotated, Literal from unittest import mock @@ -160,14 +161,36 @@ def _fake_run_loop( mock_run_loop.assert_called_once() -class FakeSolverConfig(solver_pydantic_model.LinearThetaMethod): +class FakeSolverConfig(solver_pydantic_model.BaseSolver): """Fake solver config that allows us to hook into the error logic.""" - solver_type: Annotated[Literal['fake'], torax_pydantic.JAX_STATIC] = 'fake' # pyrefly: ignore[bad-override] + solver_type: Annotated[Literal['fake'], torax_pydantic.JAX_STATIC] = 'fake' param: Annotated[str, torax_pydantic.JAX_STATIC] = 'T_i_right_bc' max_value: float = 2.5 inner_solver_iterations: list[int] | None = None + @functools.cached_property + def build_runtime_params( + self, + ) -> solver_pydantic_model.runtime_params.RuntimeParams: + return solver_pydantic_model.runtime_params.RuntimeParams( + theta_implicit=self.theta_implicit, + convection_dirichlet_mode=self.convection_dirichlet_mode, + convection_neumann_mode=self.convection_neumann_mode, + use_pereverzev=self.use_pereverzev, + use_predictor_corrector=self.use_predictor_corrector, + implicit_solver_type=self.implicit_solver_type, + chi_pereverzev=self.chi_pereverzev, + D_pereverzev=self.D_pereverzev, + n_corrector_steps=self.n_corrector_steps, + fixed_point_atol=self.fixed_point_atol, + fixed_point_rtol=self.fixed_point_rtol, + fixed_point_termination_criterion=self.fixed_point_termination_criterion, + fixed_point_sufficient_decrease=self.fixed_point_sufficient_decrease, + fixed_point_use_backtracking=self.fixed_point_use_backtracking, + delta_reduction_factor=self.delta_reduction_factor, + ) + def build_solver( self, models: models_lib.Models,