diff --git a/torax/_src/edge/extended_lengyel/extended_lengyel_model.py b/torax/_src/edge/extended_lengyel/extended_lengyel_model.py index f3518d7d5..8eacc0ae2 100644 --- a/torax/_src/edge/extended_lengyel/extended_lengyel_model.py +++ b/torax/_src/edge/extended_lengyel/extended_lengyel_model.py @@ -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) 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 5018a10eb..149e021f2 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 @@ -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 @@ -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', diff --git a/torax/_src/edge/tests/edge_updaters_test.py b/torax/_src/edge/tests/edge_updaters_test.py index c7f9030f9..fee929bf0 100644 --- a/torax/_src/edge/tests/edge_updaters_test.py +++ b/torax/_src/edge/tests/edge_updaters_test.py @@ -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() diff --git a/torax/_src/edge/updaters.py b/torax/_src/edge/updaters.py index 1dab509b2..5fa8fe9ca 100644 --- a/torax/_src/edge/updaters.py +++ b/torax/_src/edge/updaters.py @@ -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] @@ -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,