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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 9 additions & 32 deletions torax/_src/fvm/newton_raphson_solve_block.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
from torax._src import state as state_module
from torax._src.config import runtime_params as runtime_params_lib
from torax._src.core_profiles import convertors
from torax._src.fvm import block_1d_coeffs
from torax._src.fvm import calc_coeffs
from torax._src.fvm import cell_variable
from torax._src.fvm import enums
Expand Down Expand Up @@ -56,9 +57,7 @@
)
def newton_raphson_solve_block(
dt: array_typing.FloatScalar,
runtime_params_t: runtime_params_lib.RuntimeParams,
runtime_params_t_plus_dt: runtime_params_lib.RuntimeParams,
geo_t: geometry.Geometry,
geo_t_plus_dt: geometry.Geometry,
x_old: tuple[cell_variable.CellVariable, ...],
core_profiles_t: state_module.CoreProfiles,
Expand All @@ -74,6 +73,8 @@ def newton_raphson_solve_block(
delta_reduction_factor: float,
tau_min: float,
pedestal_transition_state: pedestal_transition_state_lib.PedestalTransitionState,
coeffs_old: block_1d_coeffs.Block1DCoeffs,
coeffs_exp_linear: block_1d_coeffs.Block1DCoeffs | None = None,
log_iterations: bool = False,
) -> tuple[
tuple[cell_variable.CellVariable, ...],
Expand Down Expand Up @@ -108,9 +109,7 @@ def newton_raphson_solve_block(

Args:
dt: Discrete time step.
runtime_params_t: Runtime parameters for time t.
runtime_params_t_plus_dt: Runtime parameters for time t + dt.
geo_t: Geometry at time t.
geo_t_plus_dt: Geometry at time t + dt.
x_old: Tuple containing CellVariables for each channel with their values at
the start of the time step.
Expand Down Expand Up @@ -145,6 +144,9 @@ def newton_raphson_solve_block(
routine resets at a lower timestep.
pedestal_transition_state: State for tracking pedestal L-H and H-L
transitions.
coeffs_old: PDE coefficients at beginning of timestep.
coeffs_exp_linear: PDE coefficients at beginning of timestep with additional
pereverzev terms.
log_iterations: If true, output diagnostic information from within iteration
loop.

Expand All @@ -157,38 +159,13 @@ def newton_raphson_solve_block(
"""
# pyformat: enable

coeffs_old = coeffs_callback(
runtime_params_t,
geo_t,
core_profiles_t,
prev_core_profiles=None,
dt=None,
x=x_old,
explicit_source_profiles=explicit_source_profiles,
explicit_call=True,
pedestal_transition_state=pedestal_transition_state,
)

match initial_guess_mode:
# LINEAR initial guess will provide the initial guess using the predictor-
# corrector method if predictor_corrector=True in the solver config
case enums.InitialGuessMode.LINEAR:
# returns transport coefficients with additional pereverzev terms
# if set by runtime_params, needed if stiff transport models
# (e.g. qlknn) are used.
coeffs_exp_linear = coeffs_callback(
runtime_params_t,
geo_t,
core_profiles=core_profiles_t,
prev_core_profiles=None,
dt=None,
x=x_old,
explicit_source_profiles=explicit_source_profiles,
allow_pereverzev=True,
explicit_call=True,
pedestal_transition_state=pedestal_transition_state,
)

assert (
coeffs_exp_linear is not None
), 'coeffs_exp_linear must be provided for LINEAR guess mode'
# See linear_theta_method.py for comments on the predictor_corrector API
x_new_guess = convertors.core_profiles_to_solver_x_tuple(
core_profiles_t_plus_dt, evolving_names
Expand Down
41 changes: 8 additions & 33 deletions torax/_src/fvm/optimizer_solve_block.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,9 +49,7 @@
)
def optimizer_solve_block(
dt: jax.Array,
runtime_params_t: runtime_params_lib.RuntimeParams,
runtime_params_t_plus_dt: runtime_params_lib.RuntimeParams,
geo_t: geometry.Geometry,
geo_t_plus_dt: geometry.Geometry,
x_old: tuple[cell_variable.CellVariable, ...],
core_profiles_t: state.CoreProfiles,
Expand All @@ -64,6 +62,8 @@ def optimizer_solve_block(
maxiter: int,
tol: float,
pedestal_transition_state: pedestal_transition_state_lib.PedestalTransitionState,
coeffs_old: block_1d_coeffs.Block1DCoeffs,
coeffs_exp_linear: block_1d_coeffs.Block1DCoeffs | None = None,
) -> tuple[
tuple[cell_variable.CellVariable, ...],
state.SolverNumericOutputs,
Expand All @@ -80,11 +80,7 @@ def optimizer_solve_block(

Args:
dt: Discrete time step.
runtime_params_t: Runtime params for time t (the start time of the step).
These runtime params can change from step to step without triggering a
recompilation.
runtime_params_t_plus_dt: Runtime params for time t + dt.
geo_t: Geometry object used to initialize auxiliary outputs at time t.
geo_t_plus_dt: Geometry object used to initialize auxiliary outputs at time
t + dt.
x_old: Tuple containing CellVariables for each channel with their values at
Expand Down Expand Up @@ -113,6 +109,9 @@ def optimizer_solve_block(
tol: See docstring of `jaxopt.LBFGS`.
pedestal_transition_state: State for tracking pedestal L-H and H-L
transitions.
coeffs_old: PDE coefficients at beginning of timestep.
coeffs_exp_linear: PDE coefficients at beginning of timestep with additional
pereverzev terms.

Returns:
x_new: Tuple, with x_new[i] giving channel i of x at the next time step
Expand All @@ -121,38 +120,14 @@ def optimizer_solve_block(
"""
# pyformat: enable

coeffs_old = coeffs_callback(
runtime_params_t,
geo_t,
core_profiles_t,
prev_core_profiles=None,
dt=None,
x=x_old,
explicit_source_profiles=explicit_source_profiles,
explicit_call=True,
pedestal_transition_state=pedestal_transition_state,
)

match initial_guess_mode:
# LINEAR initial guess will provide the initial guess using the predictor-
# corrector method if use_predictor_corrector=True in the solver runtime
# params
case enums.InitialGuessMode.LINEAR:
# returns transport coefficients with additional pereverzev terms
# if set by runtime_params, needed if stiff transport models (e.g. qlknn)
# are used.
coeffs_exp_linear = coeffs_callback(
runtime_params_t,
geo_t,
core_profiles_t,
prev_core_profiles=None,
dt=None,
x=x_old,
explicit_source_profiles=explicit_source_profiles,
allow_pereverzev=True,
explicit_call=True,
pedestal_transition_state=pedestal_transition_state,
)
assert (
coeffs_exp_linear is not None
), 'coeffs_exp_linear must be provided for LINEAR guess mode'
# See linear_theta_method.py for comments on the predictor_corrector API
x_new_guess = convertors.core_profiles_to_solver_x_tuple(
core_profiles_t_plus_dt, evolving_names
Expand Down
105 changes: 57 additions & 48 deletions torax/_src/mhd/sawtooth/sawtooth_solver.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@

import jax
from jax import numpy as jnp
from torax._src import array_typing
from torax._src import constants
from torax._src import jax_utils
from torax._src import state
Expand All @@ -30,6 +31,15 @@
from torax._src.sources import source_profiles as source_profiles_lib


@jax.tree_util.register_dataclass
@dataclasses.dataclass(frozen=True)
class SawtoothPreparedStepState(solver.PreparedStepState):
trigger_sawtooth: array_typing.BoolScalar
rho_norm_q1: array_typing.FloatScalar
runtime_params_t: runtime_params_lib.RuntimeParams
geo_t: geometry.Geometry


# TODO(b/414537757). Sawtooth extensions.
# a. Full and incomplete Kadomtsev redistribution model.
# b. Porcelli model with free parameters and fast ion sensitivities.
Expand All @@ -38,51 +48,20 @@
class SawtoothSolver(solver.Solver):
"""Sawtooth trigger and redistribution, and carries out sawtooth step."""

def _x_new(
@jax.jit(static_argnames=['self'])
def prepare_step(
self,
dt: jax.Array,
t: jax.Array,
runtime_params_t: runtime_params_lib.RuntimeParams,
runtime_params_t_plus_dt: runtime_params_lib.RuntimeParams,
geo_t: geometry.Geometry,
geo_t_plus_dt: geometry.Geometry,
core_profiles_t: state.CoreProfiles,
core_profiles_t_plus_dt: state.CoreProfiles,
explicit_source_profiles: source_profiles_lib.SourceProfiles,
evolving_names: tuple[str, ...],
pedestal_transition_state: (
pedestal_transition_state_lib.PedestalTransitionState
),
) -> tuple[
tuple[cell_variable.CellVariable, ...],
state.SolverNumericOutputs,
]:
"""Applies the sawtooth model and outputs new state attributes if triggered.

If the trigger model indicates a crash has been triggered, an
instantaneous redistribution model is applied. New state attributes
following a short (configurable) dt are returned. Beyond the sawtooth
redistribution, core_profiles are further updated by: the psidot assumed at
time t; the new boundary conditions; sources, transport, geo, and pedestal
outputs consistent with the new core_profiles at time t_plus_crash_dt.

Args:
dt: Sawtooth step duration.
runtime_params_t: Runtime parameters at time t.
runtime_params_t_plus_dt: Runtime parameters at time t + crash_dt.
geo_t: Geometry at time t.
geo_t_plus_dt: Geometry at time t + crash_dt.
core_profiles_t: Core profiles at time t.
core_profiles_t_plus_dt: Core profiles containing boundary conditions and
prescribed profiles at time t + crash_dt.
explicit_source_profiles: Explicit source profiles at time t.
evolving_names: Names of evolving variables.
pedestal_transition_state: State for tracking pedestal L-H and H-L
transitions.

Returns:
Updated tuple of evolving CellVariables from CoreProfiles
SolverNumericOutputs indicating a sawtooth crash.
"""
pedestal_transition_state: pedestal_transition_state_lib.PedestalTransitionState,
) -> SawtoothPreparedStepState:
evolving_names = runtime_params_t.numerics.evolving_names
x_old = convertors.core_profiles_to_solver_x_tuple(
core_profiles_t, evolving_names
)
sawtooth_models = self.models.mhd_models.sawtooth_models
if sawtooth_models is None:
raise ValueError('Sawtooth model is None.')
Expand All @@ -97,16 +76,43 @@ def _x_new(
# encounter division-by-zero.
rho_norm_q1 = jnp.maximum(rho_norm_q1, constants.CONSTANTS.eps)

return SawtoothPreparedStepState(
x_old=x_old,
core_profiles_t=core_profiles_t,
explicit_source_profiles=explicit_source_profiles,
pedestal_transition_state=pedestal_transition_state,
trigger_sawtooth=trigger_sawtooth,
rho_norm_q1=rho_norm_q1,
runtime_params_t=runtime_params_t,
geo_t=geo_t,
)

@jax.jit(static_argnames=['self'])
def solve_step(
self,
prepared_state: SawtoothPreparedStepState,
dt: jax.Array,
runtime_params_t_plus_dt: runtime_params_lib.RuntimeParams,
geo_t_plus_dt: geometry.Geometry,
core_profiles_t_plus_dt: state.CoreProfiles,
) -> tuple[
tuple[cell_variable.CellVariable, ...],
state.SolverNumericOutputs,
]:
evolving_names = runtime_params_t_plus_dt.numerics.evolving_names
sawtooth_models = self.models.mhd_models.sawtooth_models
if sawtooth_models is None:
raise ValueError('Sawtooth model is None.')

def _redistribute_state() -> tuple[
tuple[cell_variable.CellVariable, ...],
state.SolverNumericOutputs,
]:

redistributed_core_profiles = sawtooth_models.redistribution_model(
rho_norm_q1,
runtime_params_t,
geo_t,
core_profiles_t,
prepared_state.rho_norm_q1,
prepared_state.runtime_params_t,
prepared_state.geo_t,
prepared_state.core_profiles_t,
)

# Evolve the psi profile over the sawtooth time.
Expand All @@ -119,7 +125,7 @@ def _redistribute_state() -> tuple[
# using `updaters.update_all_core_profiles_after_step`.
evolved_psi_redistributed_value = (
redistributed_core_profiles.psi.value
+ core_profiles_t.psidot.value * dt
+ prepared_state.core_profiles_t.psidot.value * dt
)
evolved_core_profiles = dataclasses.replace(
redistributed_core_profiles,
Expand Down Expand Up @@ -148,10 +154,13 @@ def _redistribute_state() -> tuple[
# Return redistributed state attributes if triggered, otherwise return
# unchanged state attributes.
return jax.lax.cond(
trigger_sawtooth,
prepared_state.trigger_sawtooth,
_redistribute_state,
lambda: (
tuple([getattr(core_profiles_t, name) for name in evolving_names]),
tuple([
getattr(prepared_state.core_profiles_t, name)
for name in evolving_names
]),
state.SolverNumericOutputs(
sawtooth_crash=False,
solver_error_state=jnp.array(0, jax_utils.get_int_dtype()),
Expand Down
Loading
Loading