Skip to content
Merged
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
16 changes: 12 additions & 4 deletions torax/_src/edge/extended_lengyel/extended_lengyel_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -255,10 +255,18 @@ def __call__(
continue

# Calculate edge concentration: c_edge = c_core_lcfs * enrichment_factor
# Enrichment factor exists for all species (validated in config)
fixed_impurity_concentrations[species] = (
ratio_face[-1] * edge_params.enrichment_factor[species]
)
if (
edge_params.use_enrichment_model
and previous_edge_outputs is not None
):
assert isinstance(
previous_edge_outputs,
extended_lengyel_standalone.ExtendedLengyelOutputs,
)
enrichment = previous_edge_outputs.calculated_enrichment[species]
else:
enrichment = edge_params.enrichment_factor[species]
fixed_impurity_concentrations[species] = ratio_face[-1] * enrichment

# Determine initial guesses
initial_guess = _get_initial_guess(edge_params, previous_edge_outputs)
Expand Down
150 changes: 150 additions & 0 deletions torax/_src/edge/extended_lengyel/tests/extended_lengyel_model_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
from torax._src.edge.extended_lengyel import extended_lengyel_standalone
from torax._src.edge.extended_lengyel import pydantic_model
from torax._src.fvm import cell_variable
from torax._src.geometry import chease
from torax._src.geometry import geometry
from torax._src.geometry import standard_geometry
from torax._src.neoclassical.bootstrap_current import base as bootstrap_current_base
Expand Down Expand Up @@ -627,6 +628,155 @@ def test_initial_guess_skips_bad_previous_outputs(
err_msg='kappa_e should have fallen back to default, not bad previous.',
)

@parameterized.named_parameters(
(
'core_sot_model_on_with_previous_outputs',
extended_lengyel_model.FixedImpuritySourceOfTruth.CORE,
True,
True,
),
(
'core_sot_model_on_without_previous_outputs',
extended_lengyel_model.FixedImpuritySourceOfTruth.CORE,
True,
False,
),
(
'core_sot_model_off',
extended_lengyel_model.FixedImpuritySourceOfTruth.CORE,
False,
True,
),
(
'edge_sot_model_on',
extended_lengyel_model.FixedImpuritySourceOfTruth.EDGE,
True,
True,
),
(
'edge_sot_model_off',
extended_lengyel_model.FixedImpuritySourceOfTruth.EDGE,
False,
True,
),
)
@mock.patch.object(
extended_lengyel_standalone, 'run_extended_lengyel_standalone'
)
def test_fixed_impurity_concentrations_conditionals(
self,
impurity_sot,
use_enrichment_model,
provide_previous_outputs,
mock_run_standalone,
):
n_rho = 10
geo = chease.CheaseConfig(n_rho=n_rho).build_geometry()

mock_core_profiles = mock.MagicMock(spec=state.CoreProfiles)
for attr in ['n_e', 'n_i', 'n_impurity']:
m = mock.MagicMock(spec=cell_variable.CellVariable)
m.face_value.return_value = np.ones(n_rho + 1) * 1e19
setattr(mock_core_profiles, attr, m)
mock_psi = mock.MagicMock(spec=cell_variable.CellVariable)
mock_psi.face_value.return_value = np.linspace(0.0, 1.0, n_rho + 1)
mock_core_profiles.psi = mock_psi
mock_core_profiles.Z_i_face = np.ones(n_rho + 1)
mock_core_profiles.A_i = np.array(2.0)
mock_core_profiles.A_impurity_face = np.ones(n_rho + 1) * 20.0
mock_core_profiles.Ip_profile_face = np.ones(n_rho + 1) * 1e6

mock_core_sources = mock.MagicMock(spec=source_profiles.SourceProfiles)
mock_core_sources.total_sources.return_value = np.ones(n_rho) * 1e6

core_lcfs_ratio = 0.02
impurity_params = mock.MagicMock(spec=electron_density_ratios.RuntimeParams)
impurity_params.n_e_ratios = {'Ne': np.array(core_lcfs_ratio)}
impurity_params.n_e_ratios_face = {
'Ne': np.ones(n_rho + 1) * core_lcfs_ratio
}

previous_enrichment = 3.5
if provide_previous_outputs:
previous_edge_outputs = extended_lengyel_standalone.ExtendedLengyelOutputs(
T_e_right_bc=jnp.array(0.1),
T_i_right_bc=jnp.array(0.1),
q_parallel=jnp.array(1e8),
q_perpendicular_target=jnp.array(1e6),
T_e_separatrix=jnp.array(0.1),
T_e_target=jnp.array(3.5),
pressure_neutral_divertor=jnp.array(1.0),
alpha_t=jnp.array(0.42),
kappa_e=jnp.array(2500.0),
c_z_prefactor=jnp.array(0.0),
Z_eff_separatrix=jnp.array(1.5),
seed_impurity_concentrations={},
calculated_enrichment={'Ne': jnp.array(previous_enrichment)},
solver_status=extended_lengyel_solvers.ExtendedLengyelSolverStatus(
physics_outcome=extended_lengyel_solvers.PhysicsOutcome.SUCCESS,
numerics_outcome=(
extended_lengyel_solvers.FixedPointOutcome.SUCCESS
),
),
)
else:
previous_edge_outputs = None

edge_config = pydantic_model.ExtendedLengyelConfig(
model_name='extended_lengyel',
computation_mode=extended_lengyel_enums.ComputationMode.FORWARD,
impurity_sot=impurity_sot,
use_enrichment_model=use_enrichment_model,
enrichment_factor={'Ne': 2.0},
fixed_impurity_concentrations={'Ne': 0.05},
seed_impurity_weights={},
connection_length_target=10.0,
connection_length_divertor=2.0,
toroidal_flux_expansion=1.0,
angle_of_incidence_target=1.0,
ratio_bpol_omp_to_bpol_avg=1.0,
diverted=True,
)
runtime_params = mock.MagicMock(spec=runtime_params_lib.RuntimeParams)
runtime_params.edge = edge_config.build_runtime_params(t=0.0)
mock_pc = mock.MagicMock(spec=plasma_composition_lib.RuntimeParams)
mock_pc.impurity = impurity_params
runtime_params.plasma_composition = mock_pc

model = extended_lengyel_model.ExtendedLengyelModel()
model(
runtime_params=runtime_params,
geo=geo,
core_profiles=mock_core_profiles,
core_sources=mock_core_sources,
previous_edge_outputs=previous_edge_outputs,
)

_, kwargs = mock_run_standalone.call_args
if (
impurity_sot == extended_lengyel_model.FixedImpuritySourceOfTruth.CORE
and use_enrichment_model
and provide_previous_outputs
):
expected_concentration = core_lcfs_ratio * previous_enrichment
self.assertNotEqual(
previous_enrichment, runtime_params.edge.enrichment_factor['Ne']
)
elif impurity_sot == extended_lengyel_model.FixedImpuritySourceOfTruth.CORE:
expected_concentration = (
core_lcfs_ratio * runtime_params.edge.enrichment_factor['Ne']
)
else:
expected_concentration = (
runtime_params.edge.fixed_impurity_concentrations['Ne']
)

np.testing.assert_allclose(
kwargs['fixed_impurity_concentrations']['Ne'],
expected_concentration,
rtol=1e-5,
)

@parameterized.named_parameters(
(
'diverted',
Expand Down
93 changes: 0 additions & 93 deletions torax/_src/edge/tests/edge_updaters_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -408,99 +408,6 @@ def test_update_impurities_scales_profile_with_enrichment_model(self):
updated_n_e_ratios, initial_n_e_ratios * scaling_factor, rtol=1e-5
)

@parameterized.named_parameters(
('core_sot_model_on', 'core', True),
('edge_sot_model_on', 'edge', False),
)
def test_updates_enrichment_factor_conditionally_when_use_enrichment_model_true(
self, impurity_sot, should_update
):
self.config_dict['edge']['impurity_sot'] = impurity_sot
self.config_dict['edge']['use_enrichment_model'] = True

torax_config = model_config.ToraxConfig.from_dict(self.config_dict)
provider = build_runtime_params.RuntimeParamsProvider.from_config(
torax_config
)
runtime_params = provider(t=0.0)
assert isinstance(runtime_params.edge, extended_lengyel_model.RuntimeParams)
initial_enrichment_factor = runtime_params.edge.enrichment_factor

edge_outputs = mock.MagicMock(
spec=extended_lengyel_standalone.ExtendedLengyelOutputs
)
edge_outputs.calculated_enrichment = {
'N': jnp.array(self._CALCULATED_ENRICHMENT)
}
edge_outputs.seed_impurity_concentrations = {}
edge_outputs.T_e_right_bc = 1.0
edge_outputs.T_i_right_bc = 1.0

updated_runtime_params = updaters.update_runtime_params(
runtime_params, edge_outputs
)
assert isinstance(
updated_runtime_params.edge, extended_lengyel_model.RuntimeParams
)
updated_enrichment_factor = updated_runtime_params.edge.enrichment_factor

if should_update:
# It should be updated to the value from edge_outputs
np.testing.assert_allclose(
updated_enrichment_factor['N'], self._CALCULATED_ENRICHMENT
)
# And it should be different from the initial value
self.assertNotEqual(
initial_enrichment_factor['N'], updated_enrichment_factor['N']
)
else:
# It should not have been updated
np.testing.assert_allclose(
updated_enrichment_factor['N'], initial_enrichment_factor['N']
)

@parameterized.named_parameters(
('core_sot_model_off', 'core'),
('edge_sot_model_off', 'edge'),
)
def test_does_not_update_enrichment_factor_when_use_enrichment_model_false(
self, impurity_sot
):
self.config_dict['edge']['impurity_sot'] = impurity_sot
self.config_dict['edge']['use_enrichment_model'] = False
self.config_dict['edge']['enrichment_factor'] = {'N': 5.0}

torax_config = model_config.ToraxConfig.from_dict(self.config_dict)
provider = build_runtime_params.RuntimeParamsProvider.from_config(
torax_config
)
runtime_params = provider(t=0.0)
assert isinstance(runtime_params.edge, extended_lengyel_model.RuntimeParams)
initial_enrichment_factor = runtime_params.edge.enrichment_factor

edge_outputs = mock.MagicMock(
spec=extended_lengyel_standalone.ExtendedLengyelOutputs
)
edge_outputs.calculated_enrichment = {
'N': jnp.array(self._CALCULATED_ENRICHMENT)
}
edge_outputs.seed_impurity_concentrations = {}
edge_outputs.T_e_right_bc = 1.0
edge_outputs.T_i_right_bc = 1.0

updated_runtime_params = updaters.update_runtime_params(
runtime_params, edge_outputs
)
assert isinstance(
updated_runtime_params.edge, extended_lengyel_model.RuntimeParams
)
updated_enrichment_factor = updated_runtime_params.edge.enrichment_factor

# It should not have been updated
np.testing.assert_allclose(
updated_enrichment_factor['N'], initial_enrichment_factor['N']
)


if __name__ == '__main__':
absltest.main()
37 changes: 0 additions & 37 deletions torax/_src/edge/updaters.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,17 +52,6 @@ def update_runtime_params(

assert isinstance(runtime_params.edge, extended_lengyel_model.RuntimeParams)

if (
runtime_params.edge.use_enrichment_model
and runtime_params.edge.impurity_sot
== extended_lengyel_model.FixedImpuritySourceOfTruth.CORE
):
# If the enrichment model is used and the core is the source of truth for
# fixed impurities, then we need to update the enrichment factors in the
# runtime_params for use in the edge model, consistent with last
# edge model outputs.
runtime_params = _update_enrichment_factor(runtime_params, edge_outputs)

# Conditionally update temperatures based on the update_temperatures flag.
runtime_params = jax.lax.cond(
runtime_params.edge.update_temperatures, # pyrefly: ignore[missing-attribute]
Expand All @@ -82,32 +71,6 @@ def update_runtime_params(
return runtime_params


def _update_enrichment_factor(
runtime_params: runtime_params_lib.RuntimeParams,
edge_outputs: edge_base.EdgeModelOutputs,
) -> runtime_params_lib.RuntimeParams:
"""Updates enrichment factors based on edge model outputs."""
if not isinstance(runtime_params.edge, extended_lengyel_model.RuntimeParams):
raise ValueError(
'Enrichment factor updates from the edge model are only supported for'
' the extended Lengyel model.'
)
if not isinstance(
edge_outputs, extended_lengyel_standalone.ExtendedLengyelOutputs
):
raise ValueError(
'Enrichment factor updates from the edge model are only supported for'
' the extended Lengyel model.'
)
enrichment_factor = edge_outputs.calculated_enrichment
return dataclasses.replace(
runtime_params,
edge=dataclasses.replace(
runtime_params.edge, enrichment_factor=enrichment_factor
),
)


def _update_temperatures(
runtime_params: runtime_params_lib.RuntimeParams,
edge_outputs: edge_base.EdgeModelOutputs,
Expand Down
Loading