76 replacements: dict[str, str],
77 symbolic_functions: dict[str, SymbolicFunctionDefinition] |
None =
None,
78 expression_replacements: dict[object, str] |
None =
None,
80 symbolic_functions = symbolic_functions
or {}
81 expression_replacements = expression_replacements
or {}
82 placeholder_replacements: dict[str, str] = {}
84 def render_subexpression(subexpression) -> str:
85 for source_expression, target_expression
in expression_replacements.items():
86 if subexpression == source_expression:
87 return target_expression
92 expression_replacements,
95 def placeholder(rendered_expression: str) -> Symbol:
96 symbol = Symbol(f
"__similie_symbolic_function_{len(placeholder_replacements)}")
97 placeholder_replacements[str(symbol)] = rendered_expression
100 for function_name, function_definition
in symbolic_functions.items():
101 expression = expression.replace(
102 lambda node, name=function_name: (
103 isinstance(node, Subs)
104 and isinstance(node.expr, Derivative)
105 and node.expr.expr.func.__name__ == name
107 lambda node, definition=function_definition: placeholder(
108 definition.derivative_expressions[node.expr.derivative_count].format(
109 argument=render_subexpression(node.point[0])
113 expression = expression.replace(
114 lambda node, name=function_name: node.func.__name__ == name,
115 lambda node, definition=function_definition: placeholder(
116 definition.value_expression.format(
117 argument=render_subexpression(node.args[0])
122 for source_expression, target_expression
in expression_replacements.items():
123 expression = expression.replace(
124 lambda node, source=source_expression: node == source,
125 lambda node, target=target_expression: placeholder(target),
130 placeholder_replacements,
350 replacements: dict[str, str],
351 symbolic_functions: dict[str, SymbolicFunctionDefinition],
352 expression_replacements: dict[object, str],
353 component_expression: str,
355 row_component = component_expression.format(index=
"RowIndex", moments=
"moments")
356 column_component = component_expression.format(
357 index=
"ColumnIndex", moments=
"moments"
359 diagonal_expressions = [
361 diff(diff(hamiltonian, symbol), symbol),
362 {**replacements, str(symbol): row_component},
364 expression_replacements,
366 for symbol
in symbols_
368 off_diagonal_expressions = [
370 diff(diff(hamiltonian, symbols_[0]), symbol),
373 str(symbols_[0]): row_component,
374 str(symbol): column_component,
377 expression_replacements,
379 for symbol
in symbols_[1:]
383 "moments jacobian cannot be generated without tag-independent expressions"
387 template <class RowIndex, class ColumnIndex, class Moments, class Elem>
388 KOKKOS_FUNCTION constexpr double jacobian(Moments moments, Elem elem) const
390 static_cast<void>(elem);
391 if constexpr (std::is_same_v<RowIndex, ColumnIndex>) {{
392 return {diagonal_expressions[0]};
394 return {off_diagonal_expressions[0]};
404 parameters: list[tuple[str, str, bool, str]],
406 variable_entries: list[dict[str, object]],
407 includes: list[str] |
None =
None,
408 moments_computer: str |
None =
None,
409 inverse_symbols: list[str] |
None =
None,
410 inverse_expressions: list |
None =
None,
411 template_parameters: list[str] |
None =
None,
412 parameter_value_expressions: dict[str, str] |
None =
None,
413 definition: HamiltonianDefinition |
None =
None,
415 output_path.parent.mkdir(parents=
True, exist_ok=
True)
417 parameter_value_expressions = parameter_value_expressions
or {}
418 symbolic_functions = (
419 {}
if definition
is None else (definition.symbolic_functions
or {})
421 parameter_replacements = {}
422 for member_name, constructor_name, _, _
in parameters:
423 parameter_replacements[constructor_name] = parameter_value_expressions.get(
424 constructor_name, member_name
426 h_replacements = dict(parameter_replacements)
428 "elem" in replacement
for replacement
in parameter_replacements.values()
430 argument_signature_parts: list[str] = []
431 arguments_call_parts: list[str] = []
432 for entry
in variable_entries:
433 entry_name = entry[
"name"]
434 entry_symbols = entry[
"symbols"]
435 if len(entry_symbols) == 1:
436 argument_signature_parts.append(f
"double {entry_name}")
437 arguments_call_parts.append(entry_name)
438 h_replacements[str(entry_symbols[0])] = entry_name
440 argument_signature_parts.append(
441 f
"std::span<double const, {len(entry_symbols)}> {entry_name}"
443 arguments_call_parts.append(entry_name)
444 for i, symbol
in enumerate(entry_symbols):
445 h_replacements[str(symbol)] = f
"{entry_name}[{i}]"
449 f
" template <class Elem>\n"
450 f
" KOKKOS_FUNCTION constexpr double hamiltonian({', '.join(argument_signature_parts)}, Elem elem) const"
453 h_signature = f
" KOKKOS_FUNCTION constexpr double hamiltonian({', '.join(argument_signature_parts)}) const"
455 potential_entry = variable_entries[0]
456 moments_entry = variable_entries[1]
457 potential_symbols = potential_entry[
"symbols"]
458 moments_symbols = moments_entry[
"symbols"]
460 potential_replacements = dict(h_replacements)
461 moments_replacements = dict(parameter_replacements)
463 potential_derivative_expressions = [
464 diff(hamiltonian, symbol)
for symbol
in potential_symbols
466 moments_derivative_expressions = [
467 diff(hamiltonian, symbol)
for symbol
in moments_symbols
470 potential_argument_entries = [potential_entry, *variable_entries[2:]]
471 potential_signature_parts: list[str] = []
472 for entry
in potential_argument_entries:
473 entry_name = entry[
"name"]
474 entry_symbols = entry[
"symbols"]
475 if len(entry_symbols) == 1:
476 potential_signature_parts.append(f
"double {entry_name}")
478 potential_signature_parts.append(
479 f
"std::span<double const, {len(entry_symbols)}> {entry_name}"
482 if len(potential_symbols) == 1:
483 potential_method_signature =
", ".join(potential_signature_parts)
485 potential_method = f
"""
486 template <class Elem>
487 KOKKOS_FUNCTION constexpr double dhamiltonian_dpotential({potential_method_signature}, Elem elem) const
489 return {_render_cxx_expression(potential_derivative_expressions[0], potential_replacements, symbolic_functions)};
493 potential_method = f
"""
494 KOKKOS_FUNCTION constexpr double dhamiltonian_dpotential({potential_method_signature}) const
496 return {_render_cxx_expression(potential_derivative_expressions[0], potential_replacements, symbolic_functions)};
500 potential_method =
""
503 if inverse_symbols
is not None and inverse_expressions
is not None:
509 potential_replacements,
514 use_moments_object = (
515 len(moments_symbols) > 1
516 and definition
is not None
517 and definition.moments_object_component_expression
is not None
520 if not use_moments_object
and (
521 len(moments_symbols) == 1
or moments_computer
is None
524 "dhamiltonian_dmoments",
526 [str(symbol)
for symbol
in moments_symbols],
527 moments_derivative_expressions,
528 moments_replacements,
534 generic_moments_method =
""
535 if len(moments_symbols) > 1
and moments_computer
is not None:
536 generic_moments_replacements = {
537 **parameter_replacements,
538 **{str(symbol):
"moments" for symbol
in moments_symbols},
542 "dhamiltonian_dmoments",
544 [str(symbol)
for symbol
in moments_symbols],
545 moments_derivative_expressions,
546 generic_moments_replacements,
552 "dhamiltonian_dmoments",
554 [str(symbol)
for symbol
in moments_symbols],
555 moments_derivative_expressions,
556 generic_moments_replacements,
561 moments_object_method =
""
562 moments_jacobian_method =
""
563 if use_moments_object:
564 object_expression_replacements = {}
565 if definition.moments_object_norm2_expression
is not None:
566 object_expression_replacements[
567 sum(symbol**2
for symbol
in moments_symbols)
568 ] = definition.moments_object_norm2_expression.format(moments=
"moments")
570 [str(symbol)
for symbol
in moments_symbols],
571 moments_derivative_expressions,
572 moments_replacements,
574 object_expression_replacements,
575 definition.moments_object_component_expression,
577 if definition.generate_moments_jacobian:
581 moments_replacements,
583 object_expression_replacements,
584 definition.moments_object_component_expression,
587 nonlocal_value_methods =
""
588 if moments_computer
is not None:
590 "dhamiltonian_dmoments_value",
591 [str(symbol)
for symbol
in moments_symbols],
592 moments_derivative_expressions,
593 moments_replacements,
596 rendered_includes =
""
598 rendered_includes +=
"".join(f
"#include {header}\n" for header
in includes)
600 rendered_moments_computer =
""
601 if moments_computer
is not None:
602 rendered_moments_computer = (
603 f
" using MomentsComputer = {moments_computer};\n\n"
606 if template_parameters:
607 template_prefix =
"template <" +
", ".join(template_parameters) +
">\n"
609 is_linear =
False if definition
is None else definition.is_linear
611 output_path.write_text(
613// SPDX-FileCopyrightText: 2026 Baptiste Legouix
614// SPDX-License-Identifier: AGPL-3.0-or-later
621#include <Kokkos_Core.hpp>
623#include <type_traits>
627namespace {namespace} {{
629{"" if definition is None else definition.namespace_preamble}
631{template_prefix}struct {struct_name} {{
632 static constexpr std::size_t N = {len(moments_symbols)};
633 static constexpr bool IS_LINEAR = {"true" if is_linear else "false"};
635{rendered_moments_computer}\
636{_render_members(parameters)}
638{_render_constructor_signature(struct_name, parameters)}
639 : {_render_constructor_initializers(parameters)} {{}}
643 return {_render_cxx_expression(hamiltonian, h_replacements, symbolic_functions)};
645{potential_method}{moments_method}{generic_moments_method}{moments_object_method}{moments_jacobian_method}{nonlocal_value_methods}
648{"" if definition is None else definition.namespace_epilogue}
649}} // namespace {namespace}
655 definition = functor_class.__call__(*args, **kwargs)
656 parameter_types = definition.parameter_types
or {}
658 (f
"m_{name}", name,
True, parameter_types.get(name,
"double"))
659 for name
in definition.parameters
662 inverse_symbols =
None
663 inverse_expressions =
None
666 for entry
in definition.variables
670 if variable_entries[0][
"name"] ==
"phi":
671 dphi_dx_symbols = symbols(f
"dphi_dx0:{len(variable_entries[1]['symbols'])}")
672 inverse_solution = solve(
675 - diff(definition.hamiltonian, variable_entries[1][
"symbols"][i])
676 for i
in range(len(variable_entries[1][
"symbols"]))
678 variable_entries[1][
"symbols"],
682 inverse_symbols = [str(symbol)
for symbol
in dphi_dx_symbols]
683 inverse_expressions = [
684 inverse_solution[0][variable_entries[1][
"symbols"][i]]
685 for i
in range(len(variable_entries[1][
"symbols"]))
689 output_path=output_path,
690 namespace=definition.namespace,
691 struct_name=definition.struct_name,
692 parameters=parameter_tuples,
693 hamiltonian=definition.hamiltonian,
694 variable_entries=variable_entries,
695 includes=definition.includes,
696 moments_computer=definition.moments_computer,
697 inverse_symbols=inverse_symbols,
698 inverse_expressions=inverse_expressions,
699 template_parameters=definition.template_parameters,
700 parameter_value_expressions=definition.parameter_value_expressions,
701 definition=definition,
None write_cpp_hamiltonian_header(Path output_path, str namespace, str struct_name, list[tuple[str, str, bool, str]] parameters, hamiltonian, list[dict[str, object]] variable_entries, list[str]|None includes=None, str|None moments_computer=None, list[str]|None inverse_symbols=None, list|None inverse_expressions=None, list[str]|None template_parameters=None, dict[str, str]|None parameter_value_expressions=None, HamiltonianDefinition|None definition=None)