Source code for gkx.geometry.differentiable

"""Differentiable geometry bridge contracts for VMEC/JAX pipelines."""

from __future__ import annotations

from collections.abc import Iterator
from contextlib import contextmanager
from functools import wraps
from typing import Any

from gkx.geometry import FluxTubeGeometryData
import gkx.geometry.booz_xform_bridge as _booz_bridge
import gkx.geometry.vmec_boozer_core as _vmec_boozer_core
import gkx.geometry.vmec_flux_tube_reports as _vmec_flux_tube_reports
import gkx.geometry.vmec_state_sensitivity as _vmec_state_sensitivity
import gkx.geometry.vmec_tensor_mapping as _vmec_tensor_mapping
from gkx.geometry.autodiff_checks import (
    _json_ready as _json_ready,
    _sensitivity_conditioning_metadata,
    finite_difference_jacobian,
    observable_gradient_validation_report,
)
from gkx.geometry.backend_discovery import (
    _candidate_paths as _candidate_paths,
    _find_importable_module as _find_importable_module,
    _is_traced as _is_traced,
    discover_differentiable_geometry_backends,
)
from gkx.geometry.booz_xform_bridge import (
    evaluate_boozer_bmag_on_field_line,
)
from gkx.geometry.flux_tube_contract import (
    _ARRAY_FIELDS as _ARRAY_FIELDS,
    _GEOMETRY_OBSERVABLE_NAMES as _GEOMETRY_OBSERVABLE_NAMES,
    _array as _array,
    _scalar as _scalar,
    flux_tube_geometry_from_mapping,
    flux_tube_geometry_observables,
    geometry_observable_names,
    vmec_field_line_tensor_observable_names,
    vmec_metric_tensor_observable_names,
)
from gkx.geometry.numerics import (
    _array_parity_metrics as _array_parity_metrics,
    _boozer_half_mesh_s_grid as _boozer_half_mesh_s_grid,
    _cumulative_trapezoid as _cumulative_trapezoid,
    _evaluate_boozer_cosine_series_on_field_line,
    _interp_equal_arc_profile as _interp_equal_arc_profile,
    _interp_radial as _interp_radial,
    _periodic_bilinear_sample_2d as _periodic_bilinear_sample_2d,
    _radial_derivative_array as _radial_derivative_array,
    _radial_derivative_profile as _radial_derivative_profile,
    _scalar_parity_metrics as _scalar_parity_metrics,
)
from gkx.geometry.sensitivity import (
    geometry_inverse_design_report,
    geometry_sensitivity_report,
)


_VMEC_BOOZER_PARITY_MIN_MODE_COUNT = 21
_DEFAULT_DISCOVER_DIFFERENTIABLE_GEOMETRY_BACKENDS = (
    discover_differentiable_geometry_backends
)


@contextmanager
def _patched_module_attrs(
    module: Any, replacements: dict[str, Any]
) -> Iterator[None]:
    """Temporarily patch module attributes and restore them after the call."""

    originals = {name: getattr(module, name) for name in replacements}
    for name, value in replacements.items():
        setattr(module, name, value)
    try:
        yield
    finally:
        for name, original in originals.items():
            setattr(module, name, original)


def _call_with_facade_backend_discovery(func: Any, *args: Any, **kwargs: Any) -> Any:
    if (
        discover_differentiable_geometry_backends
        is _DEFAULT_DISCOVER_DIFFERENTIABLE_GEOMETRY_BACKENDS
    ):
        return func(*args, **kwargs)
    with _patched_module_attrs(
        _booz_bridge,
        {
            "discover_differentiable_geometry_backends": (
                discover_differentiable_geometry_backends
            )
        },
    ):
        return func(*args, **kwargs)


@wraps(_booz_bridge.vmec_boundary_aspect_sensitivity_report)
def vmec_boundary_aspect_sensitivity_report(*args: Any, **kwargs: Any) -> Any:
    return _call_with_facade_backend_discovery(
        _booz_bridge.vmec_boundary_aspect_sensitivity_report, *args, **kwargs
    )


@wraps(_booz_bridge.booz_xform_spectral_sensitivity_report)
def booz_xform_spectral_sensitivity_report(*args: Any, **kwargs: Any) -> Any:
    return _call_with_facade_backend_discovery(
        _booz_bridge.booz_xform_spectral_sensitivity_report, *args, **kwargs
    )


@wraps(_booz_bridge.booz_xform_flux_tube_mapping_from_inputs)
def booz_xform_flux_tube_mapping_from_inputs(*args: Any, **kwargs: Any) -> Any:
    return _call_with_facade_backend_discovery(
        _booz_bridge.booz_xform_flux_tube_mapping_from_inputs, *args, **kwargs
    )


@wraps(_booz_bridge.booz_xform_flux_tube_sensitivity_report)
def booz_xform_flux_tube_sensitivity_report(*args: Any, **kwargs: Any) -> Any:
    return _call_with_facade_backend_discovery(
        _booz_bridge.booz_xform_flux_tube_sensitivity_report, *args, **kwargs
    )


def _call_with_vmec_state_facade_hooks(func: Any, *args: Any, **kwargs: Any) -> Any:
    with _patched_module_attrs(
        _vmec_state_sensitivity,
        {
            "discover_differentiable_geometry_backends": (
                discover_differentiable_geometry_backends
            ),
            "booz_xform_flux_tube_mapping_from_inputs": (
                booz_xform_flux_tube_mapping_from_inputs
            ),
            "geometry_sensitivity_report": geometry_sensitivity_report,
            "finite_difference_jacobian": finite_difference_jacobian,
            "_sensitivity_conditioning_metadata": (
                _sensitivity_conditioning_metadata
            ),
        },
    ):
        return func(*args, **kwargs)


@wraps(_vmec_state_sensitivity.vmex_boozer_flux_tube_sensitivity_report)
def vmex_boozer_flux_tube_sensitivity_report(*args: Any, **kwargs: Any) -> Any:
    return _call_with_vmec_state_facade_hooks(
        _vmec_state_sensitivity.vmex_boozer_flux_tube_sensitivity_report,
        *args,
        **kwargs,
    )


@wraps(_vmec_state_sensitivity.vmex_metric_tensor_sensitivity_report)
def vmex_metric_tensor_sensitivity_report(*args: Any, **kwargs: Any) -> Any:
    return _call_with_vmec_state_facade_hooks(
        _vmec_state_sensitivity.vmex_metric_tensor_sensitivity_report,
        *args,
        **kwargs,
    )


@wraps(_vmec_state_sensitivity.vmex_field_line_tensor_sensitivity_report)
def vmex_field_line_tensor_sensitivity_report(*args: Any, **kwargs: Any) -> Any:
    return _call_with_vmec_state_facade_hooks(
        _vmec_state_sensitivity.vmex_field_line_tensor_sensitivity_report,
        *args,
        **kwargs,
    )


@wraps(_vmec_tensor_mapping.vmex_flux_tube_mapping_from_state)
def vmex_flux_tube_mapping_from_state(*args: Any, **kwargs: Any) -> Any:
    return _vmec_tensor_mapping.vmex_flux_tube_mapping_from_state(
        *args, **kwargs
    )


_cached_booz_xform_constants = _vmec_boozer_core._cached_booz_xform_constants


def _call_with_vmec_boozer_core_facade_hooks(
    func: Any, *args: Any, **kwargs: Any
) -> Any:
    with _patched_module_attrs(
        _vmec_boozer_core,
        {
            "_boozer_half_mesh_s_grid": _boozer_half_mesh_s_grid,
            "_cumulative_trapezoid": _cumulative_trapezoid,
            "_evaluate_boozer_cosine_series_on_field_line": (
                _evaluate_boozer_cosine_series_on_field_line
            ),
            "_interp_equal_arc_profile": _interp_equal_arc_profile,
            "_interp_radial": _interp_radial,
            "_radial_derivative_array": _radial_derivative_array,
            "_radial_derivative_profile": _radial_derivative_profile,
        },
    ):
        return func(*args, **kwargs)


@wraps(_vmec_boozer_core.prewarm_vmec_boozer_equal_arc_cache)
def prewarm_vmec_boozer_equal_arc_cache(*args: Any, **kwargs: Any) -> Any:
    return _call_with_vmec_boozer_core_facade_hooks(
        _vmec_boozer_core.prewarm_vmec_boozer_equal_arc_cache, *args, **kwargs
    )


@wraps(_vmec_boozer_core.vmex_boozer_equal_arc_core_profiles_from_state)
def vmex_boozer_equal_arc_core_profiles_from_state(
    *args: Any, **kwargs: Any
) -> Any:
    return _call_with_vmec_boozer_core_facade_hooks(
        _vmec_boozer_core.vmex_boozer_equal_arc_core_profiles_from_state,
        *args,
        **kwargs,
    )


[docs] def flux_tube_geometry_from_vmec_boozer_state( # pragma: no cover state: Any, runtime: Any, inp: Any, wout: Any, *, surface_index: int | None = None, torflux: float | None = None, alpha: float = 0.0, ntheta: int = 32, mboz: int = _VMEC_BOOZER_PARITY_MIN_MODE_COUNT, nboz: int = _VMEC_BOOZER_PARITY_MIN_MODE_COUNT, jit: bool = False, surface_stencil_width: int | None = None, reference_length: float | None = None, reference_b: float | None = None, source_model: str = "mode21_vmec_boozer_state", validate_finite: bool = True, ) -> FluxTubeGeometryData: """Build solver-ready geometry directly from a solved ``vmex`` state. This is the production-facing in-memory bridge for differentiable optimization workflows. It keeps the path inside JAX-compatible objects: ``SpectralState -> boozer_input_tables -> booz_xform_jax -> FluxTubeGeometryData``. Runtime VMEC file generation can still use the NetCDF/EIK route, but differentiable stellarator optimization should call this function or a higher-level objective wrapper around it so gradients never pass through filesystem artifacts. """ mapping = vmex_boozer_equal_arc_core_profiles_from_state( state, runtime, inp, wout, surface_index=surface_index, torflux=torflux, alpha=alpha, ntheta=ntheta, mboz=mboz, nboz=nboz, jit=jit, surface_stencil_width=surface_stencil_width, reference_length=reference_length, reference_b=reference_b, ) return flux_tube_geometry_from_mapping( mapping, source_model=source_model, validate_finite=validate_finite, )
def _call_with_vmec_report_facade_hooks(func: Any, *args: Any, **kwargs: Any) -> Any: with _patched_module_attrs( _vmec_flux_tube_reports, { "discover_differentiable_geometry_backends": ( discover_differentiable_geometry_backends ), "flux_tube_geometry_from_mapping": flux_tube_geometry_from_mapping, "geometry_sensitivity_report": geometry_sensitivity_report, "vmex_boozer_equal_arc_core_profiles_from_state": ( vmex_boozer_equal_arc_core_profiles_from_state ), "vmex_flux_tube_mapping_from_state": ( vmex_flux_tube_mapping_from_state ), "_array_parity_metrics": _array_parity_metrics, "_scalar_parity_metrics": _scalar_parity_metrics, }, ): return func(*args, **kwargs) @wraps(_vmec_flux_tube_reports.vmex_flux_tube_sensitivity_report) def vmex_flux_tube_sensitivity_report(*args: Any, **kwargs: Any) -> Any: return _call_with_vmec_report_facade_hooks( _vmec_flux_tube_reports.vmex_flux_tube_sensitivity_report, *args, **kwargs, ) @wraps(_vmec_flux_tube_reports.vmex_flux_tube_array_parity_report) def vmex_flux_tube_array_parity_report(*args: Any, **kwargs: Any) -> Any: return _call_with_vmec_report_facade_hooks( _vmec_flux_tube_reports.vmex_flux_tube_array_parity_report, *args, **kwargs, ) __all__ = [ "booz_xform_flux_tube_mapping_from_inputs", "booz_xform_flux_tube_sensitivity_report", "booz_xform_spectral_sensitivity_report", "discover_differentiable_geometry_backends", "evaluate_boozer_bmag_on_field_line", "finite_difference_jacobian", "flux_tube_geometry_from_mapping", "flux_tube_geometry_from_vmec_boozer_state", "flux_tube_geometry_observables", "geometry_inverse_design_report", "geometry_observable_names", "geometry_sensitivity_report", "observable_gradient_validation_report", "vmex_boozer_flux_tube_sensitivity_report", "vmex_boozer_equal_arc_core_profiles_from_state", "vmex_field_line_tensor_sensitivity_report", "vmex_flux_tube_array_parity_report", "vmex_flux_tube_mapping_from_state", "vmex_flux_tube_sensitivity_report", "vmex_metric_tensor_sensitivity_report", "vmec_boundary_aspect_sensitivity_report", "vmec_field_line_tensor_observable_names", "vmec_metric_tensor_observable_names", ]