diff --git a/madgraph/iolibs/template_files/madmatrix/umami.cc b/madgraph/iolibs/template_files/madmatrix/umami.cc index d19c93bb9..1f511a85f 100644 --- a/madgraph/iolibs/template_files/madmatrix/umami.cc +++ b/madgraph/iolibs/template_files/madmatrix/umami.cc @@ -320,6 +320,12 @@ extern "C" case UMAMI_IN_DIAGRAM_INDEX: diagram_in = static_cast( input ); break; + case UMAMI_IN_INVARIANT_COUNT: + case UMAMI_IN_INVARIANT_PIDS_AND_MASKS: + case UMAMI_IN_INVARIANT_MASSES: + case UMAMI_IN_INVARIANT_VIRTUALITIES: + // accept, but ignore externally supplied invariants + break; default: return UMAMI_ERROR_UNSUPPORTED_INPUT; } diff --git a/madgraph/iolibs/template_files/mg7/madevent.py b/madgraph/iolibs/template_files/mg7/madevent.py index 48078090d..eadf5e535 100644 --- a/madgraph/iolibs/template_files/mg7/madevent.py +++ b/madgraph/iolibs/template_files/mg7/madevent.py @@ -1060,6 +1060,9 @@ def build_multichannel_phasespace(self) -> PhaseSpace: invariant_power=self.process.run_card["phasespace"]["invariant_power"], permutations=chan_permutations, leptonic=self.process.leptonic, + return_invariants=self.process.run_card["phasespace"][ + "pass_invariants_to_matrix_element" + ], ) prefix = f"subproc{self.subproc_id}.channel{channel_id}" if topo_count > 1: @@ -1114,6 +1117,9 @@ def build_flat_phasespace(self) -> PhaseSpace: mode=self.t_channel_mode(self.process.run_card["phasespace"]["flat_mode"]), cuts=self.cuts, leptonic=self.process.leptonic, + return_invariants=self.process.run_card["phasespace"][ + "pass_invariants_to_matrix_element" + ], ) prefix = f"subproc{self.subproc_id}.flat" discrete_before, discrete_after = self.build_discrete( @@ -1377,37 +1383,62 @@ def build_integrands( flavor_remap.append(flav["index"]) flavor_factors.append(len(flav["options"])) flavor_mirror.append(flav["mirror"]) - if self.matrix_element: - matrix_element = ms.MatrixElement( - self.matrix_element, - ms.Integrand.matrix_element_inputs, - ms.Integrand.matrix_element_outputs, - True, - ) - else: - matrix_element = ms.MatrixElement( - 0xBADCAFE, - self.particle_count, - ms.Integrand.matrix_element_inputs, - ms.Integrand.matrix_element_outputs, - self.meta["diagram_count"], - True, - ) + pass_invariants = self.process.run_card["phasespace"][ + "pass_invariants_to_matrix_element" + ] + matrix_element_inputs = list(ms.Integrand.matrix_element_inputs) + if pass_invariants: + matrix_element_inputs += [ + ms.MatrixElement.invariant_count_in, + ms.MatrixElement.invariant_pids_and_masks_in, + ms.MatrixElement.invariant_masses_in, + ms.MatrixElement.invariant_virtualities_in, + ] + pdf_grid = None if self.process.leptonic else self.process.pdf_grid pdf_arg = None if self.process.leptonic else ms.CachedPdf() - cross_section = ms.DifferentialCrossSection( - matrix_element=matrix_element, - cm_energy=self.process.e_cm, - running_coupling=None, - energy_scale=ms.CachedScale(), - pid_options=flavors, - pdf1=pdf_arg, - pdf2=pdf_arg, - input_momentum_fraction=True, - ) + + def build_cross_section(invariant_count: int) -> ms.DifferentialCrossSection: + if self.matrix_element: + matrix_element = ms.MatrixElement( + self.matrix_element, + matrix_element_inputs, + ms.Integrand.matrix_element_outputs, + True, + invariant_count, + ) + else: + matrix_element = ms.MatrixElement( + 0xBADCAFE, + self.particle_count, + matrix_element_inputs, + ms.Integrand.matrix_element_outputs, + self.meta["diagram_count"], + True, + invariant_count, + ) + return ms.DifferentialCrossSection( + matrix_element=matrix_element, + cm_energy=self.process.e_cm, + running_coupling=None, + energy_scale=ms.CachedScale(), + pid_options=flavors, + pdf1=pdf_arg, + pdf2=pdf_arg, + input_momentum_fraction=True, + ) + + # only build per-integrand cross sections if invariants are passed + shared_cross_section = None if pass_invariants else build_cross_section(0) + partial_weights = self.process.run_card["generation"]["systematics"] integrands = [] for channel in phasespace.channels: + cross_section = ( + shared_cross_section + if shared_cross_section is not None + else build_cross_section(channel.phasespace_mapping.invariant_count()) + ) integrands.append(ms.Integrand( channel.phasespace_mapping, cross_section, diff --git a/madgraph/iolibs/template_files/mg7/run_card.toml b/madgraph/iolibs/template_files/mg7/run_card.toml index c48d95a2c..09c0ae5b7 100644 --- a/madgraph/iolibs/template_files/mg7/run_card.toml +++ b/madgraph/iolibs/template_files/mg7/run_card.toml @@ -64,6 +64,7 @@ flat_mode = %(phasespace.flat_mode)s # options: propagator, rambo, chili simplified_channel_count = %(phasespace.simplified_channel_count)s invariant_power = %(phasespace.invariant_power)s bw_cutoff = %(phasespace.bw_cutoff)s +pass_invariants_to_matrix_element = %(phasespace.pass_invariants_to_matrix_element)s [multiparticles] $multiparticles diff --git a/madgraph/various/RunCardLO_to_MG7_mapping.md b/madgraph/various/RunCardLO_to_MG7_mapping.md index 132ad31a8..88e780c86 100644 --- a/madgraph/various/RunCardLO_to_MG7_mapping.md +++ b/madgraph/various/RunCardLO_to_MG7_mapping.md @@ -130,7 +130,7 @@ Cuts that are **not representable** in the current MG7 cut engine ([x] unless no `generation.freeze_max_weight_after`, `generation.max_overweight_truncation`, `generation.cut_efficiency_threshold`, `generation.max_cut_repetitions`, all of `[vegas]`, `phasespace.{mode,t_channel,flat_mode,invariant_power, -simplified_channel_count,decays}`, all of `[madnis]`. +simplified_channel_count,decays,pass_invariants_to_matrix_element}`, all of `[madnis]`. --- diff --git a/madgraph/various/banner.py b/madgraph/various/banner.py index 2120ed298..54734b0e5 100755 --- a/madgraph/various/banner.py +++ b/madgraph/various/banner.py @@ -6505,6 +6505,7 @@ def default_setup(self): self.add_toml_param('phasespace', 'simplified_channel_count', 10) self.add_toml_param('phasespace', 'invariant_power', 0.7) self.add_toml_param('phasespace', 'bw_cutoff', 15) + self.add_toml_param('phasespace', 'pass_invariants_to_matrix_element', False) # ----------------------------- [madnis] ----------------------- self.add_toml_param('madnis', 'enable', False) diff --git a/madspace/include/madspace/compgraphs/function_builder_mixin.inc b/madspace/include/madspace/compgraphs/function_builder_mixin.inc index 360d8db14..fcb74c1fa 100644 --- a/madspace/include/madspace/compgraphs/function_builder_mixin.inc +++ b/madspace/include/madspace/compgraphs/function_builder_mixin.inc @@ -71,6 +71,10 @@ Value sub(Value in1, Value in2) { return instruction("sub", {in1, in2})[0]; } +Value neg(Value in) { + return instruction("neg", {in})[0]; +} + Value mul(Value in1, Value in2) { return instruction("mul", {in1, in2})[0]; } @@ -389,9 +393,9 @@ std::array uniform_invariant_inverse(Value s, Value s_min, Value s_max return {output_vector[0], output_vector[1]}; } -std::array breit_wigner_invariant(Value r, Value mass, Value width, Value s_min, Value s_max) { +std::array breit_wigner_invariant(Value r, Value mass, Value width, Value s_min, Value s_max) { auto output_vector = instruction("breit_wigner_invariant", {r, mass, width, s_min, s_max}); - return {output_vector[0], output_vector[1]}; + return {output_vector[0], output_vector[1], output_vector[2]}; } std::array breit_wigner_invariant_inverse(Value s, Value mass, Value width, Value s_min, Value s_max) { @@ -399,9 +403,9 @@ std::array breit_wigner_invariant_inverse(Value s, Value mass, Value w return {output_vector[0], output_vector[1]}; } -std::array stable_invariant(Value r, Value mass, Value s_min, Value s_max) { +std::array stable_invariant(Value r, Value mass, Value s_min, Value s_max) { auto output_vector = instruction("stable_invariant", {r, mass, s_min, s_max}); - return {output_vector[0], output_vector[1]}; + return {output_vector[0], output_vector[1], output_vector[2]}; } std::array stable_invariant_inverse(Value s, Value mass, Value s_min, Value s_max) { @@ -409,9 +413,9 @@ std::array stable_invariant_inverse(Value s, Value mass, Value s_min, return {output_vector[0], output_vector[1]}; } -std::array stable_invariant_nu(Value r, Value mass, Value nu, Value s_min, Value s_max) { +std::array stable_invariant_nu(Value r, Value mass, Value nu, Value s_min, Value s_max) { auto output_vector = instruction("stable_invariant_nu", {r, mass, nu, s_min, s_max}); - return {output_vector[0], output_vector[1]}; + return {output_vector[0], output_vector[1], output_vector[2]}; } std::array stable_invariant_nu_inverse(Value s, Value mass, Value nu, Value s_min, Value s_max) { @@ -591,6 +595,10 @@ Value permute_momenta(Value momenta, Value permutations, Value index) { return instruction("permute_momenta", {momenta, permutations, index})[0]; } +Value permute_bits(Value input, Value permutations, Value index) { + return instruction("permute_bits", {input, permutations, index})[0]; +} + Value gather(Value index, Value choices) { return instruction("gather", {index, choices})[0]; } diff --git a/madspace/include/madspace/compgraphs/opcode_mixin.inc b/madspace/include/madspace/compgraphs/opcode_mixin.inc index ba664afd0..b330d5f2a 100644 --- a/madspace/include/madspace/compgraphs/opcode_mixin.inc +++ b/madspace/include/madspace/compgraphs/opcode_mixin.inc @@ -15,143 +15,145 @@ accept_norm = 13, add = 14, add_int = 15, sub = 16, -mul = 17, -div = 18, -reduce_sum = 19, -reduce_sum_vector = 20, -batch_reduce_mean = 21, -batch_reduce_mean_keepdim = 22, -reduce_product = 23, -sqrt = 24, -square = 25, -min = 26, -max = 27, -obs_sqrt_s = 28, -obs_e = 29, -obs_px = 30, -obs_py = 31, -obs_pz = 32, -obs_mass = 33, -obs_pt = 34, -obs_p_mag = 35, -obs_phi = 36, -obs_theta = 37, -obs_y = 38, -obs_y_abs = 39, -obs_eta = 40, -obs_eta_abs = 41, -obs_delta_eta = 42, -obs_delta_phi = 43, -obs_delta_r = 44, -boost_beam = 45, -boost_beam_inverse = 46, -com_p_in = 47, -r_to_x1x2 = 48, -x1x2_to_r = 49, -diff_cross_section = 50, -two_body_decay_com = 51, -two_body_decay_com_inverse = 52, -two_body_decay = 53, -two_body_decay_inverse = 54, -two_to_two_particle_scattering_com = 55, -two_to_two_particle_scattering_com_inverse = 56, -two_to_two_particle_scattering = 57, -two_to_two_particle_scattering_inverse = 58, -two_to_three_particle_scattering = 59, -two_to_three_particle_scattering_inverse = 60, -double_t_scattering = 61, -double_t_scattering_inverse = 62, -three_body_decay_com = 63, -three_body_decay_com_inverse = 64, -three_body_decay = 65, -three_body_decay_inverse = 66, -t_inv_min_max = 67, -t_inv_value_and_min_max = 68, -t_inv_min_max_cut = 69, -t_inv_value_and_min_max_cut = 70, -t1_inv_min_max_doublet = 71, -t1_inv_value_and_min_max_doublet = 72, -t2_inv_min_max_doublet = 73, -t2_inv_value_and_min_max_doublet = 74, -s23_min_max = 75, -s23_value_and_min_max = 76, -s23_min_max_cut = 77, -s23_value_and_min_max_cut = 78, -invariants_from_momenta = 79, -sde2_channel_weights = 80, -subchannel_weights = 81, -apply_subchannel_weights = 82, -pt_eta_phi_x = 83, -mirror_momenta = 84, -momenta_to_x1x2 = 85, -uniform_invariant = 86, -uniform_invariant_inverse = 87, -breit_wigner_invariant = 88, -breit_wigner_invariant_inverse = 89, -stable_invariant = 90, -stable_invariant_inverse = 91, -stable_invariant_nu = 92, -stable_invariant_nu_inverse = 93, -fast_rambo_massless = 94, -fast_rambo_massless_inverse = 95, -fast_rambo_massless_com = 96, -fast_rambo_massive = 97, -fast_rambo_massive_inverse = 98, -fast_rambo_massive_com = 99, -cut_unphysical = 100, -cut_one = 101, -cut_all = 102, -cut_any = 103, -scale_transverse_energy = 104, -scale_transverse_mass = 105, -scale_half_transverse_mass = 106, -scale_partonic_energy = 107, -chili_forward = 108, -chili_inverse = 109, -matrix_element = 110, -collect_channel_weights = 111, -interpolate_pdf = 112, -interpolate_alpha_s = 113, -matmul = 114, -relu = 115, -leaky_relu = 116, -elu = 117, -gelu = 118, -sigmoid = 119, -softplus = 120, -rqs_reshape = 121, -rqs_find_bin = 122, -rqs_forward = 123, -rqs_inverse = 124, -softmax = 125, -softmax_prior = 126, -sample_discrete = 127, -sample_discrete_inverse = 128, -sample_discrete_probs = 129, -sample_discrete_probs_inverse = 130, -discrete_histogram = 131, -permute_momenta = 132, -gather = 133, -gather_int = 134, -gather_vector = 135, -select_int = 136, -select = 137, -select_vector = 138, -argsort = 139, -quantile = 140, -one_hot = 141, -madnis_abs_weight = 142, -madnis_softclip = 143, -madnis_variance = 144, -madnis_single_channel_variance = 145, -madnis_multi_channel_variance = 146, -nonzero = 147, -batch_gather = 148, -batch_scatter = 149, -random = 150, -random_int = 151, -unweight = 152, -vegas_forward = 153, -vegas_inverse = 154, -vegas_histogram = 155, -histogram = 156 +neg = 17, +mul = 18, +div = 19, +reduce_sum = 20, +reduce_sum_vector = 21, +batch_reduce_mean = 22, +batch_reduce_mean_keepdim = 23, +reduce_product = 24, +sqrt = 25, +square = 26, +min = 27, +max = 28, +obs_sqrt_s = 29, +obs_e = 30, +obs_px = 31, +obs_py = 32, +obs_pz = 33, +obs_mass = 34, +obs_pt = 35, +obs_p_mag = 36, +obs_phi = 37, +obs_theta = 38, +obs_y = 39, +obs_y_abs = 40, +obs_eta = 41, +obs_eta_abs = 42, +obs_delta_eta = 43, +obs_delta_phi = 44, +obs_delta_r = 45, +boost_beam = 46, +boost_beam_inverse = 47, +com_p_in = 48, +r_to_x1x2 = 49, +x1x2_to_r = 50, +diff_cross_section = 51, +two_body_decay_com = 52, +two_body_decay_com_inverse = 53, +two_body_decay = 54, +two_body_decay_inverse = 55, +two_to_two_particle_scattering_com = 56, +two_to_two_particle_scattering_com_inverse = 57, +two_to_two_particle_scattering = 58, +two_to_two_particle_scattering_inverse = 59, +two_to_three_particle_scattering = 60, +two_to_three_particle_scattering_inverse = 61, +double_t_scattering = 62, +double_t_scattering_inverse = 63, +three_body_decay_com = 64, +three_body_decay_com_inverse = 65, +three_body_decay = 66, +three_body_decay_inverse = 67, +t_inv_min_max = 68, +t_inv_value_and_min_max = 69, +t_inv_min_max_cut = 70, +t_inv_value_and_min_max_cut = 71, +t1_inv_min_max_doublet = 72, +t1_inv_value_and_min_max_doublet = 73, +t2_inv_min_max_doublet = 74, +t2_inv_value_and_min_max_doublet = 75, +s23_min_max = 76, +s23_value_and_min_max = 77, +s23_min_max_cut = 78, +s23_value_and_min_max_cut = 79, +invariants_from_momenta = 80, +sde2_channel_weights = 81, +subchannel_weights = 82, +apply_subchannel_weights = 83, +pt_eta_phi_x = 84, +mirror_momenta = 85, +momenta_to_x1x2 = 86, +uniform_invariant = 87, +uniform_invariant_inverse = 88, +breit_wigner_invariant = 89, +breit_wigner_invariant_inverse = 90, +stable_invariant = 91, +stable_invariant_inverse = 92, +stable_invariant_nu = 93, +stable_invariant_nu_inverse = 94, +fast_rambo_massless = 95, +fast_rambo_massless_inverse = 96, +fast_rambo_massless_com = 97, +fast_rambo_massive = 98, +fast_rambo_massive_inverse = 99, +fast_rambo_massive_com = 100, +cut_unphysical = 101, +cut_one = 102, +cut_all = 103, +cut_any = 104, +scale_transverse_energy = 105, +scale_transverse_mass = 106, +scale_half_transverse_mass = 107, +scale_partonic_energy = 108, +chili_forward = 109, +chili_inverse = 110, +matrix_element = 111, +collect_channel_weights = 112, +interpolate_pdf = 113, +interpolate_alpha_s = 114, +matmul = 115, +relu = 116, +leaky_relu = 117, +elu = 118, +gelu = 119, +sigmoid = 120, +softplus = 121, +rqs_reshape = 122, +rqs_find_bin = 123, +rqs_forward = 124, +rqs_inverse = 125, +softmax = 126, +softmax_prior = 127, +sample_discrete = 128, +sample_discrete_inverse = 129, +sample_discrete_probs = 130, +sample_discrete_probs_inverse = 131, +discrete_histogram = 132, +permute_momenta = 133, +permute_bits = 134, +gather = 135, +gather_int = 136, +gather_vector = 137, +select_int = 138, +select = 139, +select_vector = 140, +argsort = 141, +quantile = 142, +one_hot = 143, +madnis_abs_weight = 144, +madnis_softclip = 145, +madnis_variance = 146, +madnis_single_channel_variance = 147, +madnis_multi_channel_variance = 148, +nonzero = 149, +batch_gather = 150, +batch_scatter = 151, +random = 152, +random_int = 153, +unweight = 154, +vegas_forward = 155, +vegas_inverse = 156, +vegas_histogram = 157, +histogram = 158 diff --git a/madspace/include/madspace/compgraphs/type.hpp b/madspace/include/madspace/compgraphs/type.hpp index f0d8ada0b..787e20b21 100644 --- a/madspace/include/madspace/compgraphs/type.hpp +++ b/madspace/include/madspace/compgraphs/type.hpp @@ -118,6 +118,9 @@ const Type batch_four_vec{DataType::dt_float, batch_size, {4}}; inline Type batch_float_array(int count) { return {DataType::dt_float, batch_size, {count}}; } +inline Type batch_int_array(int count) { + return {DataType::dt_int, batch_size, {count}}; +} inline Type batch_four_vec_array(int count) { return {DataType::dt_float, batch_size, {count, 4}}; } diff --git a/madspace/include/madspace/phasespace/invariants.hpp b/madspace/include/madspace/phasespace/invariants.hpp index fa274ce4d..0f3a87d89 100644 --- a/madspace/include/madspace/phasespace/invariants.hpp +++ b/madspace/include/madspace/phasespace/invariants.hpp @@ -6,7 +6,12 @@ namespace madspace { class Invariant : public Mapping { public: - Invariant(double power = 0, double mass = 0, double width = 0); + Invariant( + double power = 0, + double mass = 0, + double width = 0, + bool return_virtuality = false + ); private: Result build_forward_impl( @@ -21,6 +26,7 @@ class Invariant : public Mapping { ) const override; double _power, _mass, _width; + bool _return_virtuality; }; } // namespace madspace diff --git a/madspace/include/madspace/phasespace/matrix_element.hpp b/madspace/include/madspace/phasespace/matrix_element.hpp index 26b5a8531..03b78e2eb 100644 --- a/madspace/include/madspace/phasespace/matrix_element.hpp +++ b/madspace/include/madspace/phasespace/matrix_element.hpp @@ -16,7 +16,14 @@ class MatrixElement : public FunctionGenerator { random_diagram_in, helicity_in, channel_in, - diagram_in + diagram_in, + // number of invariants passed as invariant_pids_and_masks_in/ + // invariant_masses_in/invariant_virtualities_in; a host constant, see + // the invariant_count constructor argument, not a per-event value + invariant_count_in, + invariant_pids_and_masks_in, + invariant_masses_in, + invariant_virtualities_in }; enum MatrixElementOutput { @@ -33,13 +40,15 @@ class MatrixElement : public FunctionGenerator { const std::vector& inputs = {momenta_in}, const std::vector& outputs = {matrix_element_out}, std::size_t diagram_count = 1, - bool sample_random_inputs = false + bool sample_random_inputs = false, + std::size_t invariant_count = 0 ); MatrixElement( const MatrixElementApi& matrix_element_api, const std::vector& inputs = {momenta_in}, const std::vector& outputs = {matrix_element_out}, - bool sample_random_inputs = false + bool sample_random_inputs = false, + std::size_t invariant_count = 0 ) : MatrixElement( matrix_element_api.index(), @@ -47,11 +56,20 @@ class MatrixElement : public FunctionGenerator { inputs, outputs, matrix_element_api.diagram_count(), - sample_random_inputs + sample_random_inputs, + invariant_count ) {}; std::size_t matrix_element_index() const { return _matrix_element_index; } std::size_t diagram_count() const { return _diagram_count; } std::size_t particle_count() const { return _particle_count; } + std::size_t invariant_count() const { + for (auto input : _inputs) { + if (input == invariant_pids_and_masks_in) { + return arg_types().at("invariant_pids_and_masks").shape.at(0); + } + } + return 0; + } const std::vector& inputs() const { return _inputs; } const std::vector& outputs() const { return _outputs; } std::vector external_inputs() const; diff --git a/madspace/include/madspace/phasespace/phasespace.hpp b/madspace/include/madspace/phasespace/phasespace.hpp index b1827565d..079c2b4c9 100644 --- a/madspace/include/madspace/phasespace/phasespace.hpp +++ b/madspace/include/madspace/phasespace/phasespace.hpp @@ -25,7 +25,8 @@ class PhaseSpaceMapping : public Mapping { TChannelMode t_channel_mode = propagator, const std::optional& cuts = std::nullopt, const std::vector>& permutations = {}, - const std::optional>& color_order = std::nullopt + const std::optional>& color_order = std::nullopt, + bool return_invariants = false ); PhaseSpaceMapping( @@ -35,7 +36,8 @@ class PhaseSpaceMapping : public Mapping { double invariant_power = 0.8, TChannelMode mode = rambo, const std::optional& cuts = std::nullopt, - const std::optional>& color_order = std::nullopt + const std::optional>& color_order = std::nullopt, + bool return_invariants = false ); std::size_t random_dim() const { @@ -46,6 +48,13 @@ class PhaseSpaceMapping : public Mapping { return _topology.outgoing_masses().size() + 2; } std::size_t channel_count() const { return _permutations.size(); } + bool return_invariants() const { return _return_invariants; } + std::size_t invariant_count() const { + if (!_return_invariants) { + return 0; + } + return output_types().at("invariant_pids_and_masks").shape.at(0); + } private: Result build_forward_impl( @@ -76,6 +85,7 @@ class PhaseSpaceMapping : public Mapping { _t_mapping; std::vector> _s_decays; nested_vector2 _permutations; + bool _return_invariants; }; } // namespace madspace diff --git a/madspace/include/madspace/phasespace/t_propagator_mapping.hpp b/madspace/include/madspace/phasespace/t_propagator_mapping.hpp index 81a97423d..46881abbc 100644 --- a/madspace/include/madspace/phasespace/t_propagator_mapping.hpp +++ b/madspace/include/madspace/phasespace/t_propagator_mapping.hpp @@ -14,10 +14,21 @@ class TPropagatorMapping : public Mapping { TPropagatorMapping( const std::vector& integration_order, double invariant_power = 0.8, - const std::vector& pt_min = {} + const std::vector& pt_min = {}, + bool return_invariants = false ); std::size_t random_dim() const { return 3 * _integration_order.size() - 1; } + // Masks for every invariant sampled by TPropagatorMapping: one per t-channel + // propagator (physical order along the chain), followed by one per + // intermediate s-channel propagator. outgoing_masks are the masks for the + // attached legs + std::vector invariant_masks(const std::vector& outgoing_masks) const; + + static std::size_t invariant_count(std::size_t t_propagator_count) { + return 2 * t_propagator_count - 1; + } + private: Result build_forward_impl( FunctionBuilder& fb, @@ -37,6 +48,7 @@ class TPropagatorMapping : public Mapping { std::vector _sample_sides; std::vector _pt_min; bool _has_cut; + bool _return_invariants; Invariant _uniform_invariant; TwoToTwoParticleScattering _com_scattering; TwoToTwoParticleScattering _lab_scattering; diff --git a/madspace/include/madspace/phasespace/topology.hpp b/madspace/include/madspace/phasespace/topology.hpp index c023e970d..ca9a5ecaa 100644 --- a/madspace/include/madspace/phasespace/topology.hpp +++ b/madspace/include/madspace/phasespace/topology.hpp @@ -72,6 +72,7 @@ class Topology { double e_min; double e_max; int pdg_id; + int momentum_mask; bool on_shell; bool on_shell_boundary; }; diff --git a/madspace/include/madspace/phasespace/two_particle.hpp b/madspace/include/madspace/phasespace/two_particle.hpp index 238c6b135..df2d05dc6 100644 --- a/madspace/include/madspace/phasespace/two_particle.hpp +++ b/madspace/include/madspace/phasespace/two_particle.hpp @@ -32,7 +32,8 @@ class TwoToTwoParticleScattering : public Mapping { double invariant_power = 0, double mass = 0, double width = 0, - bool has_cut = false + bool has_cut = false, + bool return_invariant = false ); private: @@ -50,6 +51,7 @@ class TwoToTwoParticleScattering : public Mapping { bool _com; Invariant _invariant; bool _has_cut; + bool _return_invariant; }; class DoubleT : public Mapping { diff --git a/madspace/include/madspace/umami.h b/madspace/include/madspace/umami.h index 57d079d1c..e8fe85727 100644 --- a/madspace/include/madspace/umami.h +++ b/madspace/include/madspace/umami.h @@ -93,10 +93,24 @@ typedef enum { UMAMI_IN_DIAGRAM_INDEX, /** externally selected channel index, type: `unsigned int`, shape: `()` */ UMAMI_IN_CHANNEL_INDEX, + /** single integer specifying the number of invariants passed to the matrix element, + * type: `unsigned int` */ + UMAMI_IN_INVARIANT_COUNT, + /** upper 16 bits: 0 if massless invariant was sampled, otherwise PDG ID of the + * propagator. Lower 16 bits: binary mask specifying the external momenta entering + * the invariant, where the incoming momenta contribute negatively. Type: `int`, + * shape: `(invariant count)` */ + UMAMI_IN_INVARIANT_PIDS_AND_MASKS, + /** invariant masses, i.e. `E^2 - p_x^2 - p_y^2 - p_z^2`, type: `double`, + * shape: `(invariant count)` */ + UMAMI_IN_INVARIANT_MASSES, + /** virtuality, i.e. `E^2 - p_x^2 - p_y^2 - p_z^2 - m^2`, type: `double`, + * shape: `(invariant count)` */ + UMAMI_IN_INVARIANT_VIRTUALITIES, } UmamiInputKey; /** Number of values in `UmamiInputKey` */ -#define UMAMI_INPUT_KEY_COUNT (UMAMI_IN_CHANNEL_INDEX + 1) +#define UMAMI_INPUT_KEY_COUNT (UMAMI_IN_INVARIANT_VIRTUALITIES + 1) typedef enum { /** value of the matrix element, type: `double`, shape: `()` */ diff --git a/madspace/instruction_set.yaml b/madspace/instruction_set.yaml index cc8055381..cf37785e6 100644 --- a/madspace/instruction_set.yaml +++ b/madspace/instruction_set.yaml @@ -199,6 +199,18 @@ sub: differentiable: True dims: 0 +neg: + inputs: + - name: in + type: [float, ...] + desc: input + outputs: + - name: out + type: [float, ...] + desc: negative of input + desc: Negates the input. + dims: 0 + mul: inputs: - name: in1 @@ -1757,6 +1769,9 @@ breit_wigner_invariant: - name: s type: [float] desc: + - name: virt + type: [float] + desc: - name: gs type: [float] desc: @@ -1804,6 +1819,9 @@ stable_invariant: - name: s type: [float] desc: + - name: virt + type: [float] + desc: - name: gs type: [float] desc: @@ -1851,6 +1869,9 @@ stable_invariant_nu: - name: s type: [float] desc: + - name: virt + type: [float] + desc: - name: gs type: [float] desc: @@ -2536,6 +2557,22 @@ permute_momenta: type: [float, n, 4] desc: +permute_bits: + inputs: + - name: input + type: [int, k] + desc: + - name: permutations + type: [int, single, m, n] + desc: + - name: index + type: [int] + desc: + outputs: + - name: output + type: [int, k] + desc: + gather: inputs: - name: index diff --git a/madspace/src/compgraphs/instruction.cpp b/madspace/src/compgraphs/instruction.cpp index c5f6d88e9..1378ee360 100644 --- a/madspace/src/compgraphs/instruction.cpp +++ b/madspace/src/compgraphs/instruction.cpp @@ -868,6 +868,7 @@ TypeVec MatrixElementInstruction::signature(const ValueVec& args) const { case UMAMI_IN_HELICITY_INDEX: case UMAMI_IN_DIAGRAM_INDEX: case UMAMI_IN_CHANNEL_INDEX: + case UMAMI_IN_INVARIANT_COUNT: if (input_type.dtype != DataType::dt_int || input_type.shape.size() != 0) { throw std::invalid_argument( std::format( @@ -876,6 +877,26 @@ TypeVec MatrixElementInstruction::signature(const ValueVec& args) const { ); } break; + case UMAMI_IN_INVARIANT_PIDS_AND_MASKS: + if (input_type.dtype != DataType::dt_int || input_type.shape.size() != 1) { + throw std::invalid_argument( + std::format( + "matrix_element, argument {}: expected array of integers", i + 1 + ) + ); + } + break; + case UMAMI_IN_INVARIANT_MASSES: + case UMAMI_IN_INVARIANT_VIRTUALITIES: + if (input_type.dtype != DataType::dt_float || + input_type.shape.size() != 1) { + throw std::invalid_argument( + std::format( + "matrix_element, argument {}: expected array of floats", i + 1 + ) + ); + } + break; default: throw std::invalid_argument( std::format( diff --git a/madspace/src/compgraphs/instruction_set_mixin.inc b/madspace/src/compgraphs/instruction_set_mixin.inc index a2dbe6cd8..959d01e1b 100644 --- a/madspace/src/compgraphs/instruction_set_mixin.inc +++ b/madspace/src/compgraphs/instruction_set_mixin.inc @@ -32,144 +32,146 @@ InstructionOwner instructions[] { mi("add", 14, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), mi("add_int", 15, true, {{DataType::dt_int, false, {std::monostate{}}, false}, {DataType::dt_int, false, {std::monostate{}}, false}}, {{DataType::dt_int, false, {std::monostate{}}, false}}), mi("sub", 16, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("mul", 17, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("div", 18, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("reduce_sum", 19, true, {{DataType::dt_float, false, {std::monostate{}, "n"}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("reduce_sum_vector", 20, true, {{DataType::dt_float, false, {std::monostate{}, "n", "m"}, false}}, {{DataType::dt_float, false, {std::monostate{}, "m"}, false}}), - mi("batch_reduce_mean", 21, true, {{DataType::dt_float, false, {}, false}}, {{DataType::dt_float, true, {}, false}}), - mi("batch_reduce_mean_keepdim", 22, true, {{DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("reduce_product", 23, true, {{DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("sqrt", 24, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("square", 25, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("min", 26, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("max", 27, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_sqrt_s", 28, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("obs_e", 29, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_px", 30, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_py", 31, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_pz", 32, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_mass", 33, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_pt", 34, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_p_mag", 35, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_phi", 36, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_theta", 37, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_y", 38, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_y_abs", 39, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_eta", 40, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_eta_abs", 41, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_delta_eta", 42, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}, {DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_delta_phi", 43, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}, {DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("obs_delta_r", 44, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}, {DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("boost_beam", 45, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}}), - mi("boost_beam_inverse", 46, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}}), - mi("com_p_in", 47, true, {{DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}), - mi("r_to_x1x2", 48, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("x1x2_to_r", 49, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("diff_cross_section", 50, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("two_body_decay_com", 51, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("two_body_decay_com_inverse", 52, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("two_body_decay", 53, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("two_body_decay_inverse", 54, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("two_to_two_particle_scattering_com", 55, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("two_to_two_particle_scattering_com_inverse", 56, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("two_to_two_particle_scattering", 57, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("two_to_two_particle_scattering_inverse", 58, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("two_to_three_particle_scattering", 59, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("two_to_three_particle_scattering_inverse", 60, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("double_t_scattering", 61, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("double_t_scattering_inverse", 62, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("three_body_decay_com", 63, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("three_body_decay_com_inverse", 64, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("three_body_decay", 65, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("three_body_decay_inverse", 66, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("t_inv_min_max", 67, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("t_inv_value_and_min_max", 68, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("t_inv_min_max_cut", 69, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("t_inv_value_and_min_max_cut", 70, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("t1_inv_min_max_doublet", 71, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("t1_inv_value_and_min_max_doublet", 72, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("t2_inv_min_max_doublet", 73, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("t2_inv_value_and_min_max_doublet", 74, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("s23_min_max", 75, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("s23_value_and_min_max", 76, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("s23_min_max_cut", 77, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("s23_value_and_min_max_cut", 78, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("invariants_from_momenta", 79, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {"m", "n"}, false}}, {{DataType::dt_float, false, {"m"}, false}}), - mi("sde2_channel_weights", 80, true, {{DataType::dt_float, false, {"m"}, false}, {DataType::dt_float, false, {"c", "n"}, false}, {DataType::dt_float, false, {"c", "n"}, false}, {DataType::dt_int, false, {"c", "n"}, false}}, {{DataType::dt_float, false, {"c"}, false}}), - mi("subchannel_weights", 81, true, {{DataType::dt_float, false, {"m"}, false}, {DataType::dt_float, false, {"c", "n"}, false}, {DataType::dt_float, false, {"c", "n"}, false}, {DataType::dt_int, false, {"c", "n"}, false}, {DataType::dt_int, false, {"c", "n"}, false}, {DataType::dt_int, true, {"g"}, false}}, {{DataType::dt_float, false, {"c"}, false}}), - mi("apply_subchannel_weights", 82, true, {{DataType::dt_float, false, {"c"}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_int, false, {"d"}, false}, {DataType::dt_int, false, {"d"}, false}}, {{DataType::dt_float, false, {"d"}, false}}), - mi("pt_eta_phi_x", 83, true, {{DataType::dt_float, false, {"n+2", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"3n+2"}, false}}), - mi("mirror_momenta", 84, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_int, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}}), - mi("momenta_to_x1x2", 85, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("uniform_invariant", 86, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("uniform_invariant_inverse", 87, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("breit_wigner_invariant", 88, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("breit_wigner_invariant_inverse", 89, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("stable_invariant", 90, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("stable_invariant_inverse", 91, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("stable_invariant_nu", 92, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("stable_invariant_nu_inverse", 93, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("fast_rambo_massless", 94, true, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}), - mi("fast_rambo_massless_inverse", 95, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("fast_rambo_massless_com", 96, true, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}), - mi("fast_rambo_massive", 97, true, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}), - mi("fast_rambo_massive_inverse", 98, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), - mi("fast_rambo_massive_com", 99, true, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}), - mi("cut_unphysical", 100, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("cut_one", 101, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("cut_all", 102, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("cut_any", 103, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("scale_transverse_energy", 104, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("scale_transverse_mass", 105, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("scale_half_transverse_mass", 106, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("scale_partonic_energy", 107, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("chili_forward", 108, true, {{DataType::dt_float, false, {"3n-2"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"n+2", 4}, false}, {DataType::dt_float, false, {}, false}}), - mi("chili_inverse", 109, true, {{DataType::dt_float, false, {"n+2", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"3n-2"}, false}, {DataType::dt_float, false, {}, false}}), - InstructionOwner(new MatrixElementInstruction(110, true)), - mi("collect_channel_weights", 111, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_int, true, {"n"}, false}, {DataType::dt_int, true, {"c"}, true}}, {{DataType::dt_float, false, {"c"}, false}}), - mi("interpolate_pdf", 112, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, false, {"n"}, false}, {DataType::dt_float, true, {"a"}, false}, {DataType::dt_float, true, {"b"}, false}, {DataType::dt_float, true, {16, "c", "d"}, false}}, {{DataType::dt_float, false, {"n"}, false}}), - mi("interpolate_alpha_s", 113, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, true, {"b+1"}, false}, {DataType::dt_float, true, {4, "b"}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("matmul", 114, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, true, {"m", "n"}, false}, {DataType::dt_float, true, {"m"}, false}}, {{DataType::dt_float, false, {"m"}, false}}), - mi("relu", 115, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("leaky_relu", 116, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("elu", 117, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("gelu", 118, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("sigmoid", 119, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("softplus", 120, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - InstructionOwner(new RqsReshapeInstruction(121, true)), - mi("rqs_find_bin", 122, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n", "b"}, false}, {DataType::dt_float, false, {"n", "b"}, false}, {DataType::dt_float, false, {"n", "b+1"}, false}}, {{DataType::dt_float, false, {"n", 6}, false}}), - mi("rqs_forward", 123, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n", 6}, false}}, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}), - mi("rqs_inverse", 124, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n", 6}, false}}, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}), - mi("softmax", 125, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), - mi("softmax_prior", 126, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"n"}, false}}), - mi("sample_discrete", 127, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_int, false, {}, false}}, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("sample_discrete_inverse", 128, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_int, true, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("sample_discrete_probs", 129, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("sample_discrete_probs_inverse", 130, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), - mi("discrete_histogram", 131, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, true, {"n"}, true}}, {{DataType::dt_float, true, {"n"}, false}, {DataType::dt_int, true, {"n"}, false}}), - mi("permute_momenta", 132, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_int, true, {"m", "n"}, false}, {DataType::dt_int, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}}), - mi("gather", 133, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("gather_int", 134, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_int, false, {"n"}, false}}, {{DataType::dt_int, false, {}, false}}), - mi("gather_vector", 135, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {"n", "m"}, false}}, {{DataType::dt_float, false, {"m"}, false}}), - mi("select_int", 136, true, {{DataType::dt_int, false, {"n"}, false}, {DataType::dt_int, false, {"m"}, false}}, {{DataType::dt_int, false, {"m"}, false}}), - mi("select", 137, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_int, false, {"m"}, false}}, {{DataType::dt_float, false, {"m"}, false}}), - mi("select_vector", 138, true, {{DataType::dt_float, false, {"n", "k"}, false}, {DataType::dt_int, false, {"m"}, false}}, {{DataType::dt_float, false, {"m", "k"}, false}}), - mi("argsort", 139, true, {{DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_int, false, {"n"}, false}}), - mi("quantile", 140, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, true, {}, false}}, {{DataType::dt_float, true, {}, false}}), - mi("one_hot", 141, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_int, true, {"n"}, true}}, {{DataType::dt_float, false, {"n"}, false}}), - mi("madnis_abs_weight", 142, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("madnis_softclip", 143, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("madnis_variance", 144, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("madnis_single_channel_variance", 145, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), - mi("madnis_multi_channel_variance", 146, true, {{DataType::dt_float, false, {"c"}, false}, {DataType::dt_float, false, {"c"}, false}}, {{DataType::dt_float, false, {}, false}}), - InstructionOwner(new NonzeroInstruction(147, true)), - InstructionOwner(new BatchGatherInstruction(148, true)), - InstructionOwner(new BatchScatterInstruction(149, true)), - InstructionOwner(new RandomInstruction(150, true)), - InstructionOwner(new RandomIntInstruction(151, true)), - InstructionOwner(new UnweightInstruction(152, true)), - mi("vegas_forward", 153, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, true, {"n", "b"}, false}}, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}), - mi("vegas_inverse", 154, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, true, {"n", "b"}, false}}, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}), - mi("vegas_histogram", 155, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, true, {"b"}, true}}, {{DataType::dt_float, true, {"n", "b"}, false}, {DataType::dt_int, true, {"n", "b"}, false}}), - mi("histogram", 156, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, true, {"b"}, true}}, {{DataType::dt_float, true, {"b+2"}, false}, {DataType::dt_float, true, {"b+2"}, false}}), + mi("neg", 17, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("mul", 18, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("div", 19, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("reduce_sum", 20, true, {{DataType::dt_float, false, {std::monostate{}, "n"}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("reduce_sum_vector", 21, true, {{DataType::dt_float, false, {std::monostate{}, "n", "m"}, false}}, {{DataType::dt_float, false, {std::monostate{}, "m"}, false}}), + mi("batch_reduce_mean", 22, true, {{DataType::dt_float, false, {}, false}}, {{DataType::dt_float, true, {}, false}}), + mi("batch_reduce_mean_keepdim", 23, true, {{DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("reduce_product", 24, true, {{DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("sqrt", 25, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("square", 26, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("min", 27, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("max", 28, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_sqrt_s", 29, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("obs_e", 30, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_px", 31, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_py", 32, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_pz", 33, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_mass", 34, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_pt", 35, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_p_mag", 36, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_phi", 37, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_theta", 38, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_y", 39, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_y_abs", 40, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_eta", 41, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_eta_abs", 42, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_delta_eta", 43, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}, {DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_delta_phi", 44, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}, {DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_delta_r", 45, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}, {DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("boost_beam", 46, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}}), + mi("boost_beam_inverse", 47, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}}), + mi("com_p_in", 48, true, {{DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}), + mi("r_to_x1x2", 49, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("x1x2_to_r", 50, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("diff_cross_section", 51, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("two_body_decay_com", 52, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_body_decay_com_inverse", 53, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_body_decay", 54, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_body_decay_inverse", 55, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_to_two_particle_scattering_com", 56, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_to_two_particle_scattering_com_inverse", 57, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_to_two_particle_scattering", 58, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_to_two_particle_scattering_inverse", 59, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_to_three_particle_scattering", 60, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_to_three_particle_scattering_inverse", 61, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("double_t_scattering", 62, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("double_t_scattering_inverse", 63, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("three_body_decay_com", 64, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("three_body_decay_com_inverse", 65, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("three_body_decay", 66, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("three_body_decay_inverse", 67, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("t_inv_min_max", 68, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("t_inv_value_and_min_max", 69, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("t_inv_min_max_cut", 70, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("t_inv_value_and_min_max_cut", 71, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("t1_inv_min_max_doublet", 72, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("t1_inv_value_and_min_max_doublet", 73, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("t2_inv_min_max_doublet", 74, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("t2_inv_value_and_min_max_doublet", 75, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("s23_min_max", 76, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("s23_value_and_min_max", 77, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("s23_min_max_cut", 78, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("s23_value_and_min_max_cut", 79, true, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("invariants_from_momenta", 80, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {"m", "n"}, false}}, {{DataType::dt_float, false, {"m"}, false}}), + mi("sde2_channel_weights", 81, true, {{DataType::dt_float, false, {"m"}, false}, {DataType::dt_float, false, {"c", "n"}, false}, {DataType::dt_float, false, {"c", "n"}, false}, {DataType::dt_int, false, {"c", "n"}, false}}, {{DataType::dt_float, false, {"c"}, false}}), + mi("subchannel_weights", 82, true, {{DataType::dt_float, false, {"m"}, false}, {DataType::dt_float, false, {"c", "n"}, false}, {DataType::dt_float, false, {"c", "n"}, false}, {DataType::dt_int, false, {"c", "n"}, false}, {DataType::dt_int, false, {"c", "n"}, false}, {DataType::dt_int, true, {"g"}, false}}, {{DataType::dt_float, false, {"c"}, false}}), + mi("apply_subchannel_weights", 83, true, {{DataType::dt_float, false, {"c"}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_int, false, {"d"}, false}, {DataType::dt_int, false, {"d"}, false}}, {{DataType::dt_float, false, {"d"}, false}}), + mi("pt_eta_phi_x", 84, true, {{DataType::dt_float, false, {"n+2", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"3n+2"}, false}}), + mi("mirror_momenta", 85, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_int, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}}), + mi("momenta_to_x1x2", 86, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("uniform_invariant", 87, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("uniform_invariant_inverse", 88, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("breit_wigner_invariant", 89, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("breit_wigner_invariant_inverse", 90, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("stable_invariant", 91, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("stable_invariant_inverse", 92, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("stable_invariant_nu", 93, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("stable_invariant_nu_inverse", 94, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("fast_rambo_massless", 95, true, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}), + mi("fast_rambo_massless_inverse", 96, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("fast_rambo_massless_com", 97, true, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}), + mi("fast_rambo_massive", 98, true, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {4}, false}}, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}), + mi("fast_rambo_massive_inverse", 99, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {}, false}}), + mi("fast_rambo_massive_com", 100, true, {{DataType::dt_float, false, {"3n-4"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}}), + mi("cut_unphysical", 101, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("cut_one", 102, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("cut_all", 103, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("cut_any", 104, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("scale_transverse_energy", 105, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("scale_transverse_mass", 106, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("scale_half_transverse_mass", 107, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("scale_partonic_energy", 108, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("chili_forward", 109, true, {{DataType::dt_float, false, {"3n-2"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"n+2", 4}, false}, {DataType::dt_float, false, {}, false}}), + mi("chili_inverse", 110, true, {{DataType::dt_float, false, {"n+2", 4}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"3n-2"}, false}, {DataType::dt_float, false, {}, false}}), + InstructionOwner(new MatrixElementInstruction(111, true)), + mi("collect_channel_weights", 112, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_int, true, {"n"}, false}, {DataType::dt_int, true, {"c"}, true}}, {{DataType::dt_float, false, {"c"}, false}}), + mi("interpolate_pdf", 113, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, false, {"n"}, false}, {DataType::dt_float, true, {"a"}, false}, {DataType::dt_float, true, {"b"}, false}, {DataType::dt_float, true, {16, "c", "d"}, false}}, {{DataType::dt_float, false, {"n"}, false}}), + mi("interpolate_alpha_s", 114, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, true, {"b+1"}, false}, {DataType::dt_float, true, {4, "b"}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("matmul", 115, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, true, {"m", "n"}, false}, {DataType::dt_float, true, {"m"}, false}}, {{DataType::dt_float, false, {"m"}, false}}), + mi("relu", 116, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("leaky_relu", 117, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("elu", 118, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("gelu", 119, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("sigmoid", 120, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("softplus", 121, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + InstructionOwner(new RqsReshapeInstruction(122, true)), + mi("rqs_find_bin", 123, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n", "b"}, false}, {DataType::dt_float, false, {"n", "b"}, false}, {DataType::dt_float, false, {"n", "b+1"}, false}}, {{DataType::dt_float, false, {"n", 6}, false}}), + mi("rqs_forward", 124, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n", 6}, false}}, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}), + mi("rqs_inverse", 125, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n", 6}, false}}, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}), + mi("softmax", 126, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("softmax_prior", 127, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"n"}, false}}), + mi("sample_discrete", 128, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_int, false, {}, false}}, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("sample_discrete_inverse", 129, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_int, true, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("sample_discrete_probs", 130, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("sample_discrete_probs_inverse", 131, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("discrete_histogram", 132, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, true, {"n"}, true}}, {{DataType::dt_float, true, {"n"}, false}, {DataType::dt_int, true, {"n"}, false}}), + mi("permute_momenta", 133, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_int, true, {"m", "n"}, false}, {DataType::dt_int, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}}), + mi("permute_bits", 134, true, {{DataType::dt_int, false, {"k"}, false}, {DataType::dt_int, true, {"m", "n"}, false}, {DataType::dt_int, false, {}, false}}, {{DataType::dt_int, false, {"k"}, false}}), + mi("gather", 135, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("gather_int", 136, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_int, false, {"n"}, false}}, {{DataType::dt_int, false, {}, false}}), + mi("gather_vector", 137, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_float, false, {"n", "m"}, false}}, {{DataType::dt_float, false, {"m"}, false}}), + mi("select_int", 138, true, {{DataType::dt_int, false, {"n"}, false}, {DataType::dt_int, false, {"m"}, false}}, {{DataType::dt_int, false, {"m"}, false}}), + mi("select", 139, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_int, false, {"m"}, false}}, {{DataType::dt_float, false, {"m"}, false}}), + mi("select_vector", 140, true, {{DataType::dt_float, false, {"n", "k"}, false}, {DataType::dt_int, false, {"m"}, false}}, {{DataType::dt_float, false, {"m", "k"}, false}}), + mi("argsort", 141, true, {{DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_int, false, {"n"}, false}}), + mi("quantile", 142, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, true, {}, false}}, {{DataType::dt_float, true, {}, false}}), + mi("one_hot", 143, true, {{DataType::dt_int, false, {}, false}, {DataType::dt_int, true, {"n"}, true}}, {{DataType::dt_float, false, {"n"}, false}}), + mi("madnis_abs_weight", 144, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("madnis_softclip", 145, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("madnis_variance", 146, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("madnis_single_channel_variance", 147, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("madnis_multi_channel_variance", 148, true, {{DataType::dt_float, false, {"c"}, false}, {DataType::dt_float, false, {"c"}, false}}, {{DataType::dt_float, false, {}, false}}), + InstructionOwner(new NonzeroInstruction(149, true)), + InstructionOwner(new BatchGatherInstruction(150, true)), + InstructionOwner(new BatchScatterInstruction(151, true)), + InstructionOwner(new RandomInstruction(152, true)), + InstructionOwner(new RandomIntInstruction(153, true)), + InstructionOwner(new UnweightInstruction(154, true)), + mi("vegas_forward", 155, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, true, {"n", "b"}, false}}, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}), + mi("vegas_inverse", 156, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, true, {"n", "b"}, false}}, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}), + mi("vegas_histogram", 157, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, true, {"b"}, true}}, {{DataType::dt_float, true, {"n", "b"}, false}, {DataType::dt_int, true, {"n", "b"}, false}}), + mi("histogram", 158, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_int, true, {"b"}, true}}, {{DataType::dt_float, true, {"b+2"}, false}, {DataType::dt_float, true, {"b+2"}, false}}), }; diff --git a/madspace/src/cpu/runtime_backward_mixin.inc b/madspace/src/cpu/runtime_backward_mixin.inc index 8a75bf498..4120ea351 100644 --- a/madspace/src/cpu/runtime_backward_mixin.inc +++ b/madspace/src/cpu/runtime_backward_mixin.inc @@ -31,87 +31,87 @@ case 14: case 16: backward_batch_foreach, backward_kernel_sub, 3, 2, DeviceType>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 17: +case 18: backward_batch_foreach, backward_kernel_mul, 3, 2, DeviceType>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 18: +case 19: backward_batch_foreach, backward_kernel_div, 3, 2, DeviceType>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 21: +case 22: backward_op_batch_reduce_mean(instr, locals, local_grads, device); break; -case 22: +case 23: backward_op_batch_reduce_mean_keepdim(instr, locals, local_grads, device); break; -case 23: +case 24: backward_batch_foreach, backward_kernel_reduce_product, 2, 1, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 24: +case 25: backward_batch_foreach, backward_kernel_sqrt, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 25: +case 26: backward_batch_foreach, backward_kernel_square, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 114: +case 115: backward_op_matmul(instr, locals, local_grads, device); break; -case 115: +case 116: backward_batch_foreach, backward_kernel_relu, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 116: +case 117: backward_batch_foreach, backward_kernel_leaky_relu, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 117: +case 118: backward_batch_foreach, backward_kernel_elu, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 118: +case 119: backward_batch_foreach, backward_kernel_gelu, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 119: +case 120: backward_batch_foreach, backward_kernel_sigmoid, 2, 1, DeviceType>, 1, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 120: +case 121: backward_batch_foreach, backward_kernel_softplus, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 121: +case 122: backward_op_rqs_reshape(instr, locals, local_grads, device); break; -case 122: +case 123: backward_batch_foreach, backward_kernel_rqs_find_bin, 5, 4, 2, DeviceType>, 4, 1, 4, 0>(instr, locals, local_grads, {0,1,2,3}, {}, device); break; -case 123: +case 124: backward_batch_foreach, backward_kernel_rqs_forward, 4, 2, 2, DeviceType>, 2, 2, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 124: +case 125: backward_batch_foreach, backward_kernel_rqs_inverse, 4, 2, 2, DeviceType>, 2, 2, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 125: +case 126: backward_batch_foreach, backward_kernel_softmax, 2, 1, DeviceType>, 1, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 126: +case 127: backward_batch_foreach, backward_kernel_softmax_prior, 2, 2, 1, DeviceType>, 2, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 130: +case 131: backward_batch_foreach, backward_kernel_sample_discrete_probs_inverse, 4, 2, 1, DeviceType>, 2, 2, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 133: +case 135: backward_batch_foreach, backward_kernel_gather, 2, 2, 1, DeviceType>, 2, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 137: +case 139: backward_batch_foreach, backward_kernel_select, 2, 2, 1, DeviceType>, 2, 1, 1, 0>(instr, locals, local_grads, {1}, {}, device); break; -case 142: +case 144: backward_batch_foreach, backward_kernel_madnis_abs_weight, 3, 2, 1, DeviceType>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 143: +case 145: backward_batch_foreach, backward_kernel_madnis_softclip, 5, 4, 1, DeviceType>, 4, 1, 4, 0>(instr, locals, local_grads, {0,1,2,3}, {}, device); break; -case 144: +case 146: backward_batch_foreach, backward_kernel_madnis_variance, 5, 4, 1, DeviceType>, 4, 1, 4, 0>(instr, locals, local_grads, {0,1,2,3}, {}, device); break; -case 145: +case 147: backward_batch_foreach, backward_kernel_madnis_single_channel_variance, 2, 2, 1, DeviceType>, 2, 1, 1, 0>(instr, locals, local_grads, {1}, {}, device); break; -case 146: +case 148: backward_batch_foreach, backward_kernel_madnis_multi_channel_variance, 3, 2, 1, DeviceType>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; diff --git a/madspace/src/cpu/runtime_mixin.inc b/madspace/src/cpu/runtime_mixin.inc index 64a3974be..4e434594e 100644 --- a/madspace/src/cpu/runtime_mixin.inc +++ b/madspace/src/cpu/runtime_mixin.inc @@ -53,422 +53,428 @@ case 16: batch_foreach, kernel_sub, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 17: - batch_foreach, kernel_mul, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_neg, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 18: - batch_foreach, kernel_div, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_mul, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 19: - batch_foreach, kernel_reduce_sum, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_div, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 20: - batch_foreach, kernel_reduce_sum_vector, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_reduce_sum, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 21: - op_batch_reduce_mean(instr, locals, device); + batch_foreach, kernel_reduce_sum_vector, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 22: - op_batch_reduce_mean_keepdim(instr, locals, device); + op_batch_reduce_mean(instr, locals, device); break; case 23: - batch_foreach, kernel_reduce_product, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + op_batch_reduce_mean_keepdim(instr, locals, device); break; case 24: - batch_foreach, kernel_sqrt, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_reduce_product, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 25: - batch_foreach, kernel_square, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_sqrt, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 26: - batch_foreach, kernel_min, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_square, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 27: - batch_foreach, kernel_max, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_min, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 28: - batch_foreach, kernel_obs_sqrt_s, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_max, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 29: - batch_foreach, kernel_obs_e, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_sqrt_s, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 30: - batch_foreach, kernel_obs_px, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_e, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 31: - batch_foreach, kernel_obs_py, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_px, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 32: - batch_foreach, kernel_obs_pz, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_py, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 33: - batch_foreach, kernel_obs_mass, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_pz, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 34: - batch_foreach, kernel_obs_pt, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_mass, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 35: - batch_foreach, kernel_obs_p_mag, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_pt, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 36: - batch_foreach, kernel_obs_phi, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_p_mag, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 37: - batch_foreach, kernel_obs_theta, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_phi, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 38: - batch_foreach, kernel_obs_y, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_theta, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 39: - batch_foreach, kernel_obs_y_abs, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_y, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 40: - batch_foreach, kernel_obs_eta, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_y_abs, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 41: - batch_foreach, kernel_obs_eta_abs, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_eta, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 42: - batch_foreach, kernel_obs_delta_eta, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_obs_eta_abs, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 43: - batch_foreach, kernel_obs_delta_phi, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_obs_delta_eta, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 44: - batch_foreach, kernel_obs_delta_r, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_obs_delta_phi, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 45: - batch_foreach, kernel_boost_beam, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + batch_foreach, kernel_obs_delta_r, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 46: - batch_foreach, kernel_boost_beam_inverse, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + batch_foreach, kernel_boost_beam, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 47: - batch_foreach, kernel_com_p_in, 1, 2, 1, DeviceType>, 1, 2>(instr, locals, device); + batch_foreach, kernel_boost_beam_inverse, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 48: - batch_foreach, kernel_r_to_x1x2, 3, 3, 1, DeviceType>, 3, 3>(instr, locals, device); + batch_foreach, kernel_com_p_in, 1, 2, 1, DeviceType>, 1, 2>(instr, locals, device); break; case 49: - batch_foreach, kernel_x1x2_to_r, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); + batch_foreach, kernel_r_to_x1x2, 3, 3, 1, DeviceType>, 3, 3>(instr, locals, device); break; case 50: - batch_foreach, kernel_diff_cross_section, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); + batch_foreach, kernel_x1x2_to_r, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); break; case 51: - batch_foreach, kernel_two_body_decay_com, 5, 3, 1, DeviceType>, 5, 3>(instr, locals, device); + batch_foreach, kernel_diff_cross_section, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); break; case 52: - batch_foreach, kernel_two_body_decay_com_inverse, 2, 6, 1, DeviceType>, 2, 6>(instr, locals, device); + batch_foreach, kernel_two_body_decay_com, 5, 3, 1, DeviceType>, 5, 3>(instr, locals, device); break; case 53: - batch_foreach, kernel_two_body_decay, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); + batch_foreach, kernel_two_body_decay_com_inverse, 2, 6, 1, DeviceType>, 2, 6>(instr, locals, device); break; case 54: - batch_foreach, kernel_two_body_decay_inverse, 2, 7, 1, DeviceType>, 2, 7>(instr, locals, device); + batch_foreach, kernel_two_body_decay, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); break; case 55: - batch_foreach, kernel_two_to_two_particle_scattering_com, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); + batch_foreach, kernel_two_body_decay_inverse, 2, 7, 1, DeviceType>, 2, 7>(instr, locals, device); break; case 56: - batch_foreach, kernel_two_to_two_particle_scattering_com_inverse, 4, 4, 1, DeviceType>, 4, 4>(instr, locals, device); + batch_foreach, kernel_two_to_two_particle_scattering_com, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); break; case 57: - batch_foreach, kernel_two_to_two_particle_scattering, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); + batch_foreach, kernel_two_to_two_particle_scattering_com_inverse, 4, 4, 1, DeviceType>, 4, 4>(instr, locals, device); break; case 58: - batch_foreach, kernel_two_to_two_particle_scattering_inverse, 4, 4, 1, DeviceType>, 4, 4>(instr, locals, device); + batch_foreach, kernel_two_to_two_particle_scattering, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); break; case 59: - batch_foreach, kernel_two_to_three_particle_scattering, 8, 3, 1, DeviceType>, 8, 3>(instr, locals, device); + batch_foreach, kernel_two_to_two_particle_scattering_inverse, 4, 4, 1, DeviceType>, 4, 4>(instr, locals, device); break; case 60: - batch_foreach, kernel_two_to_three_particle_scattering_inverse, 7, 4, 1, DeviceType>, 7, 4>(instr, locals, device); + batch_foreach, kernel_two_to_three_particle_scattering, 8, 3, 1, DeviceType>, 8, 3>(instr, locals, device); break; case 61: - batch_foreach, kernel_double_t_scattering, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); + batch_foreach, kernel_two_to_three_particle_scattering_inverse, 7, 4, 1, DeviceType>, 7, 4>(instr, locals, device); break; case 62: - batch_foreach, kernel_double_t_scattering_inverse, 4, 2, 1, DeviceType>, 4, 2>(instr, locals, device); + batch_foreach, kernel_double_t_scattering, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); break; case 63: - batch_foreach, kernel_three_body_decay_com, 9, 4, 1, DeviceType>, 9, 4>(instr, locals, device); + batch_foreach, kernel_double_t_scattering_inverse, 4, 2, 1, DeviceType>, 4, 2>(instr, locals, device); break; case 64: - batch_foreach, kernel_three_body_decay_com_inverse, 3, 10, 1, DeviceType>, 3, 10>(instr, locals, device); + batch_foreach, kernel_three_body_decay_com, 9, 4, 1, DeviceType>, 9, 4>(instr, locals, device); break; case 65: - batch_foreach, kernel_three_body_decay, 10, 4, 1, DeviceType>, 10, 4>(instr, locals, device); + batch_foreach, kernel_three_body_decay_com_inverse, 3, 10, 1, DeviceType>, 3, 10>(instr, locals, device); break; case 66: - batch_foreach, kernel_three_body_decay_inverse, 3, 11, 1, DeviceType>, 3, 11>(instr, locals, device); + batch_foreach, kernel_three_body_decay, 10, 4, 1, DeviceType>, 10, 4>(instr, locals, device); break; case 67: - batch_foreach, kernel_t_inv_min_max, 4, 2, 1, DeviceType>, 4, 2>(instr, locals, device); + batch_foreach, kernel_three_body_decay_inverse, 3, 11, 1, DeviceType>, 3, 11>(instr, locals, device); break; case 68: - batch_foreach, kernel_t_inv_value_and_min_max, 4, 3, 1, DeviceType>, 4, 3>(instr, locals, device); + batch_foreach, kernel_t_inv_min_max, 4, 2, 1, DeviceType>, 4, 2>(instr, locals, device); break; case 69: - batch_foreach, kernel_t_inv_min_max_cut, 6, 2, 1, DeviceType>, 6, 2>(instr, locals, device); + batch_foreach, kernel_t_inv_value_and_min_max, 4, 3, 1, DeviceType>, 4, 3>(instr, locals, device); break; case 70: - batch_foreach, kernel_t_inv_value_and_min_max_cut, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); + batch_foreach, kernel_t_inv_min_max_cut, 6, 2, 1, DeviceType>, 6, 2>(instr, locals, device); break; case 71: - batch_foreach, kernel_t1_inv_min_max_doublet, 6, 2, 1, DeviceType>, 6, 2>(instr, locals, device); + batch_foreach, kernel_t_inv_value_and_min_max_cut, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); break; case 72: - batch_foreach, kernel_t1_inv_value_and_min_max_doublet, 7, 3, 1, DeviceType>, 7, 3>(instr, locals, device); + batch_foreach, kernel_t1_inv_min_max_doublet, 6, 2, 1, DeviceType>, 6, 2>(instr, locals, device); break; case 73: - batch_foreach, kernel_t2_inv_min_max_doublet, 7, 2, 1, DeviceType>, 7, 2>(instr, locals, device); + batch_foreach, kernel_t1_inv_value_and_min_max_doublet, 7, 3, 1, DeviceType>, 7, 3>(instr, locals, device); break; case 74: - batch_foreach, kernel_t2_inv_value_and_min_max_doublet, 8, 3, 1, DeviceType>, 8, 3>(instr, locals, device); + batch_foreach, kernel_t2_inv_min_max_doublet, 7, 2, 1, DeviceType>, 7, 2>(instr, locals, device); break; case 75: - batch_foreach, kernel_s23_min_max, 6, 2, 1, DeviceType>, 6, 2>(instr, locals, device); + batch_foreach, kernel_t2_inv_value_and_min_max_doublet, 8, 3, 1, DeviceType>, 8, 3>(instr, locals, device); break; case 76: - batch_foreach, kernel_s23_value_and_min_max, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); + batch_foreach, kernel_s23_min_max, 6, 2, 1, DeviceType>, 6, 2>(instr, locals, device); break; case 77: - batch_foreach, kernel_s23_min_max_cut, 10, 2, 1, DeviceType>, 10, 2>(instr, locals, device); + batch_foreach, kernel_s23_value_and_min_max, 6, 3, 1, DeviceType>, 6, 3>(instr, locals, device); break; case 78: - batch_foreach, kernel_s23_value_and_min_max_cut, 10, 3, 1, DeviceType>, 10, 3>(instr, locals, device); + batch_foreach, kernel_s23_min_max_cut, 10, 2, 1, DeviceType>, 10, 2>(instr, locals, device); break; case 79: - batch_foreach, kernel_invariants_from_momenta, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_s23_value_and_min_max_cut, 10, 3, 1, DeviceType>, 10, 3>(instr, locals, device); break; case 80: - batch_foreach, kernel_sde2_channel_weights, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); + batch_foreach, kernel_invariants_from_momenta, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 81: - batch_foreach, kernel_subchannel_weights, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); + batch_foreach, kernel_sde2_channel_weights, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); break; case 82: - batch_foreach, kernel_apply_subchannel_weights, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); + batch_foreach, kernel_subchannel_weights, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); break; case 83: - batch_foreach, kernel_pt_eta_phi_x, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + batch_foreach, kernel_apply_subchannel_weights, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); break; case 84: - batch_foreach, kernel_mirror_momenta, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_pt_eta_phi_x, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 85: - batch_foreach, kernel_momenta_to_x1x2, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); + batch_foreach, kernel_mirror_momenta, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 86: - batch_foreach, kernel_uniform_invariant, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); + batch_foreach, kernel_momenta_to_x1x2, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); break; case 87: - batch_foreach, kernel_uniform_invariant_inverse, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); + batch_foreach, kernel_uniform_invariant, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); break; case 88: - batch_foreach, kernel_breit_wigner_invariant, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); + batch_foreach, kernel_uniform_invariant_inverse, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); break; case 89: - batch_foreach, kernel_breit_wigner_invariant_inverse, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); + batch_foreach, kernel_breit_wigner_invariant, 5, 3, 1, DeviceType>, 5, 3>(instr, locals, device); break; case 90: - batch_foreach, kernel_stable_invariant, 4, 2, 1, DeviceType>, 4, 2>(instr, locals, device); + batch_foreach, kernel_breit_wigner_invariant_inverse, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); break; case 91: - batch_foreach, kernel_stable_invariant_inverse, 4, 2, 1, DeviceType>, 4, 2>(instr, locals, device); + batch_foreach, kernel_stable_invariant, 4, 3, 1, DeviceType>, 4, 3>(instr, locals, device); break; case 92: - batch_foreach, kernel_stable_invariant_nu, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); + batch_foreach, kernel_stable_invariant_inverse, 4, 2, 1, DeviceType>, 4, 2>(instr, locals, device); break; case 93: - batch_foreach, kernel_stable_invariant_nu_inverse, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); + batch_foreach, kernel_stable_invariant_nu, 5, 3, 1, DeviceType>, 5, 3>(instr, locals, device); break; case 94: - batch_foreach, kernel_fast_rambo_massless, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); + batch_foreach, kernel_stable_invariant_nu_inverse, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); break; case 95: - batch_foreach, kernel_fast_rambo_massless_inverse, 2, 3, 1, DeviceType>, 2, 3>(instr, locals, device); + batch_foreach, kernel_fast_rambo_massless, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); break; case 96: - batch_foreach, kernel_fast_rambo_massless_com, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); + batch_foreach, kernel_fast_rambo_massless_inverse, 2, 3, 1, DeviceType>, 2, 3>(instr, locals, device); break; case 97: - batch_foreach, kernel_fast_rambo_massive, 4, 2, 1, DeviceType>, 4, 2>(instr, locals, device); + batch_foreach, kernel_fast_rambo_massless_com, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); break; case 98: - batch_foreach, kernel_fast_rambo_massive_inverse, 3, 3, 1, DeviceType>, 3, 3>(instr, locals, device); + batch_foreach, kernel_fast_rambo_massive, 4, 2, 1, DeviceType>, 4, 2>(instr, locals, device); break; case 99: - batch_foreach, kernel_fast_rambo_massive_com, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); + batch_foreach, kernel_fast_rambo_massive_inverse, 3, 3, 1, DeviceType>, 3, 3>(instr, locals, device); break; case 100: - batch_foreach, kernel_cut_unphysical, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); + batch_foreach, kernel_fast_rambo_massive_com, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); break; case 101: - batch_foreach, kernel_cut_one, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + batch_foreach, kernel_cut_unphysical, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); break; case 102: - batch_foreach, kernel_cut_all, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + batch_foreach, kernel_cut_one, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 103: - batch_foreach, kernel_cut_any, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + batch_foreach, kernel_cut_all, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 104: - batch_foreach, kernel_scale_transverse_energy, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_cut_any, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 105: - batch_foreach, kernel_scale_transverse_mass, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_scale_transverse_energy, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 106: - batch_foreach, kernel_scale_half_transverse_mass, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_scale_transverse_mass, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 107: - batch_foreach, kernel_scale_partonic_energy, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_scale_half_transverse_mass, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 108: - batch_foreach, kernel_chili_forward, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); + batch_foreach, kernel_scale_partonic_energy, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 109: - batch_foreach, kernel_chili_inverse, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); + batch_foreach, kernel_chili_forward, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); break; case 110: - op_matrix_element(instr, locals, device); + batch_foreach, kernel_chili_inverse, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); break; case 111: - batch_foreach, kernel_collect_channel_weights, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + op_matrix_element(instr, locals, device); break; case 112: - batch_foreach, kernel_interpolate_pdf, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); + batch_foreach, kernel_collect_channel_weights, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 113: - batch_foreach, kernel_interpolate_alpha_s, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + batch_foreach, kernel_interpolate_pdf, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); break; case 114: - op_matmul(instr, locals, device); + batch_foreach, kernel_interpolate_alpha_s, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 115: - batch_foreach, kernel_relu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + op_matmul(instr, locals, device); break; case 116: - batch_foreach, kernel_leaky_relu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_relu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 117: - batch_foreach, kernel_elu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_leaky_relu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 118: - batch_foreach, kernel_gelu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_elu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 119: - batch_foreach, kernel_sigmoid, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_gelu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 120: - batch_foreach, kernel_softplus, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_sigmoid, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 121: - op_rqs_reshape(instr, locals, device); + batch_foreach, kernel_softplus, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 122: - batch_foreach, kernel_rqs_find_bin, 4, 1, 2, DeviceType>, 4, 1>(instr, locals, device); + op_rqs_reshape(instr, locals, device); break; case 123: - batch_foreach, kernel_rqs_forward, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); + batch_foreach, kernel_rqs_find_bin, 4, 1, 2, DeviceType>, 4, 1>(instr, locals, device); break; case 124: - batch_foreach, kernel_rqs_inverse, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); + batch_foreach, kernel_rqs_forward, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); break; case 125: - batch_foreach, kernel_softmax, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_rqs_inverse, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); break; case 126: - batch_foreach, kernel_softmax_prior, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_softmax, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 127: - batch_foreach, kernel_sample_discrete, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); + batch_foreach, kernel_softmax_prior, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 128: - batch_foreach, kernel_sample_discrete_inverse, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); + batch_foreach, kernel_sample_discrete, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); break; case 129: - batch_foreach, kernel_sample_discrete_probs, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); + batch_foreach, kernel_sample_discrete_inverse, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); break; case 130: - batch_foreach, kernel_sample_discrete_probs_inverse, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); + batch_foreach, kernel_sample_discrete_probs, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); break; case 131: - op_discrete_histogram(instr, locals, device); + batch_foreach, kernel_sample_discrete_probs_inverse, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); break; case 132: - batch_foreach, kernel_permute_momenta, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + op_discrete_histogram(instr, locals, device); break; case 133: - batch_foreach, kernel_gather, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_permute_momenta, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 134: - batch_foreach, kernel_gather_int, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_permute_bits, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 135: - batch_foreach, kernel_gather_vector, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_gather, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 136: - batch_foreach, kernel_select_int, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_gather_int, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 137: - batch_foreach, kernel_select, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_gather_vector, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 138: - batch_foreach, kernel_select_vector, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_select_int, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 139: - batch_foreach, kernel_argsort, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_select, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 140: - op_quantile(instr, locals, device); + batch_foreach, kernel_select_vector, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 141: - batch_foreach, kernel_one_hot, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_argsort, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 142: - batch_foreach, kernel_madnis_abs_weight, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + op_quantile(instr, locals, device); break; case 143: - batch_foreach, kernel_madnis_softclip, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); + batch_foreach, kernel_one_hot, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 144: - batch_foreach, kernel_madnis_variance, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); + batch_foreach, kernel_madnis_abs_weight, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 145: - batch_foreach, kernel_madnis_single_channel_variance, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_madnis_softclip, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); break; case 146: - batch_foreach, kernel_madnis_multi_channel_variance, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_madnis_variance, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); break; case 147: - op_nonzero(instr, locals, device); + batch_foreach, kernel_madnis_single_channel_variance, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 148: - op_batch_gather(instr, locals, device); + batch_foreach, kernel_madnis_multi_channel_variance, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 149: - op_batch_scatter(instr, locals, device); + op_nonzero(instr, locals, device); break; case 150: - op_random(instr, locals, device); + op_batch_gather(instr, locals, device); break; case 151: - op_random_int(instr, locals, device); + op_batch_scatter(instr, locals, device); break; case 152: - op_unweight(instr, locals, device); + op_random(instr, locals, device); break; case 153: - batch_foreach, kernel_vegas_forward, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); + op_random_int(instr, locals, device); break; case 154: - batch_foreach, kernel_vegas_inverse, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); + op_unweight(instr, locals, device); break; case 155: - op_vegas_histogram(instr, locals, device); + batch_foreach, kernel_vegas_forward, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); break; case 156: + batch_foreach, kernel_vegas_inverse, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); + break; +case 157: + op_vegas_histogram(instr, locals, device); + break; +case 158: op_histogram(instr, locals, device); break; diff --git a/madspace/src/cpu/simd_arm.hpp b/madspace/src/cpu/simd_arm.hpp index 432b99225..8ab7993b7 100644 --- a/madspace/src/cpu/simd_arm.hpp +++ b/madspace/src/cpu/simd_arm.hpp @@ -152,6 +152,15 @@ inline IVec operator-(IVec arg1) { return vnegq_s64(arg1); } inline IVec operator+(IVec arg1, IVec arg2) { return vaddq_s64(arg1, arg2); } inline IVec operator-(IVec arg1, IVec arg2) { return vsubq_s64(arg1, arg2); } +inline IVec operator&(IVec arg1, IVec arg2) { return vandq_s64(arg1, arg2); } +inline IVec operator|(IVec arg1, IVec arg2) { return vorrq_s64(arg1, arg2); } +inline IVec operator^(IVec arg1, IVec arg2) { return veorq_s64(arg1, arg2); } +inline IVec operator~(IVec arg1) { return veorq_s64(arg1, vdupq_n_s64(-1)); } +inline IVec operator<<(IVec arg1, IVec arg2) { return vshlq_s64(arg1, arg2); } +inline IVec operator>>(IVec arg1, IVec arg2) { + return vshlq_s64(arg1, vnegq_s64(arg2)); +} + inline BVec isnan(FVec arg) { return arg != arg; } inline FVec sqrt(FVec arg1) { return Sleef_sqrtd2_u05(arg1); } diff --git a/madspace/src/cpu/simd_x86_256.hpp b/madspace/src/cpu/simd_x86_256.hpp index feb8dcdd3..e045a366d 100644 --- a/madspace/src/cpu/simd_x86_256.hpp +++ b/madspace/src/cpu/simd_x86_256.hpp @@ -198,6 +198,15 @@ inline IVec operator-(IVec arg1) { inline IVec operator+(IVec arg1, IVec arg2) { return _mm256_add_epi32(arg1, arg2); } inline IVec operator-(IVec arg1, IVec arg2) { return _mm256_sub_epi32(arg1, arg2); } +inline IVec operator&(IVec arg1, IVec arg2) { return _mm256_and_si256(arg1, arg2); } +inline IVec operator|(IVec arg1, IVec arg2) { return _mm256_or_si256(arg1, arg2); } +inline IVec operator^(IVec arg1, IVec arg2) { return _mm256_xor_si256(arg1, arg2); } +inline IVec operator~(IVec arg1) { + return _mm256_xor_si256(arg1, _mm256_cmpeq_epi64(arg1, arg1)); +} +inline IVec operator<<(IVec arg1, IVec arg2) { return _mm256_sllv_epi64(arg1, arg2); } +inline IVec operator>>(IVec arg1, IVec arg2) { return _mm256_srlv_epi64(arg1, arg2); } + inline BVec isnan(FVec arg) { return arg != arg; } inline FVec sqrt(FVec arg1) { return Sleef_sqrtd4_u05avx2(arg1); } diff --git a/madspace/src/cpu/simd_x86_512.hpp b/madspace/src/cpu/simd_x86_512.hpp index c8972ab1b..8a273489d 100644 --- a/madspace/src/cpu/simd_x86_512.hpp +++ b/madspace/src/cpu/simd_x86_512.hpp @@ -199,6 +199,15 @@ inline IVec operator-(IVec arg1) { inline IVec operator+(IVec arg1, IVec arg2) { return _mm512_add_epi64(arg1, arg2); } inline IVec operator-(IVec arg1, IVec arg2) { return _mm512_sub_epi64(arg1, arg2); } +inline IVec operator&(IVec arg1, IVec arg2) { return _mm512_and_si512(arg1, arg2); } +inline IVec operator|(IVec arg1, IVec arg2) { return _mm512_or_si512(arg1, arg2); } +inline IVec operator^(IVec arg1, IVec arg2) { return _mm512_xor_si512(arg1, arg2); } +inline IVec operator~(IVec arg1) { + return _mm512_xor_si512(arg1, _mm512_set1_epi64(-1)); +} +inline IVec operator<<(IVec arg1, IVec arg2) { return _mm512_sllv_epi64(arg1, arg2); } +inline IVec operator>>(IVec arg1, IVec arg2) { return _mm512_srlv_epi64(arg1, arg2); } + inline BVec isnan(FVec arg) { return arg != arg; } inline FVec sqrt(FVec arg1) { return Sleef_sqrtd8_u05avx512f(arg1); } diff --git a/madspace/src/gpu/runtime_backward_mixin.inc b/madspace/src/gpu/runtime_backward_mixin.inc index 452946c95..6e010e7bb 100644 --- a/madspace/src/gpu/runtime_backward_mixin.inc +++ b/madspace/src/gpu/runtime_backward_mixin.inc @@ -31,87 +31,87 @@ case 14: case 16: backward_batch_foreach, 3, 2>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 17: +case 18: backward_batch_foreach, 3, 2>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 18: +case 19: backward_batch_foreach, 3, 2>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 21: +case 22: backward_op_batch_reduce_mean(instr, locals, local_grads, device); break; -case 22: +case 23: backward_op_batch_reduce_mean_keepdim(instr, locals, local_grads, device); break; -case 23: +case 24: backward_batch_foreach, 2, 1, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 24: +case 25: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 25: +case 26: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 114: +case 115: backward_op_matmul(instr, locals, local_grads, device); break; -case 115: +case 116: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 116: +case 117: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 117: +case 118: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 118: +case 119: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 119: +case 120: backward_batch_foreach, 2, 1>, 1, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 120: +case 121: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 121: +case 122: backward_op_rqs_reshape(instr, locals, local_grads, device); break; -case 122: +case 123: backward_batch_foreach, 5, 4, 2>, 4, 1, 4, 0>(instr, locals, local_grads, {0,1,2,3}, {}, device); break; -case 123: +case 124: backward_batch_foreach, 4, 2, 2>, 2, 2, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 124: +case 125: backward_batch_foreach, 4, 2, 2>, 2, 2, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 125: +case 126: backward_batch_foreach, 2, 1>, 1, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 126: +case 127: backward_batch_foreach, 2, 2, 1>, 2, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 130: +case 131: backward_batch_foreach, 4, 2, 1>, 2, 2, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 133: +case 135: backward_batch_foreach, 2, 2, 1>, 2, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 137: +case 139: backward_batch_foreach, 2, 2, 1>, 2, 1, 1, 0>(instr, locals, local_grads, {1}, {}, device); break; -case 142: +case 144: backward_batch_foreach, 3, 2, 1>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 143: +case 145: backward_batch_foreach, 5, 4, 1>, 4, 1, 4, 0>(instr, locals, local_grads, {0,1,2,3}, {}, device); break; -case 144: +case 146: backward_batch_foreach, 5, 4, 1>, 4, 1, 4, 0>(instr, locals, local_grads, {0,1,2,3}, {}, device); break; -case 145: +case 147: backward_batch_foreach, 2, 2, 1>, 2, 1, 1, 0>(instr, locals, local_grads, {1}, {}, device); break; -case 146: +case 148: backward_batch_foreach, 3, 2, 1>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; diff --git a/madspace/src/gpu/runtime_mixin.inc b/madspace/src/gpu/runtime_mixin.inc index 316c746a3..8190fcb22 100644 --- a/madspace/src/gpu/runtime_mixin.inc +++ b/madspace/src/gpu/runtime_mixin.inc @@ -53,422 +53,428 @@ case 16: batch_foreach, 2, 1>, 2, 1>(instr, locals, device); break; case 17: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 18: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1>, 2, 1>(instr, locals, device); break; case 19: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 2, 1>, 2, 1>(instr, locals, device); break; case 20: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 21: - op_batch_reduce_mean(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 22: - op_batch_reduce_mean_keepdim(instr, locals, device); + op_batch_reduce_mean(instr, locals, device); break; case 23: - batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); + op_batch_reduce_mean_keepdim(instr, locals, device); break; case 24: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); break; case 25: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 26: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 27: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1>, 2, 1>(instr, locals, device); break; case 28: - batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 2, 1>, 2, 1>(instr, locals, device); break; case 29: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); break; case 30: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 31: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 32: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 33: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 34: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 35: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 36: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 37: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 38: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 39: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 40: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 41: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 42: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 43: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1>, 2, 1>(instr, locals, device); break; case 44: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1>, 2, 1>(instr, locals, device); break; case 45: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 2, 1>, 2, 1>(instr, locals, device); break; case 46: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 47: - batch_foreach, 1, 2, 1>, 1, 2>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 48: - batch_foreach, 3, 3, 1>, 3, 3>(instr, locals, device); + batch_foreach, 1, 2, 1>, 1, 2>(instr, locals, device); break; case 49: - batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); + batch_foreach, 3, 3, 1>, 3, 3>(instr, locals, device); break; case 50: - batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); + batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); break; case 51: - batch_foreach, 5, 3, 1>, 5, 3>(instr, locals, device); + batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); break; case 52: - batch_foreach, 2, 6, 1>, 2, 6>(instr, locals, device); + batch_foreach, 5, 3, 1>, 5, 3>(instr, locals, device); break; case 53: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 2, 6, 1>, 2, 6>(instr, locals, device); break; case 54: - batch_foreach, 2, 7, 1>, 2, 7>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 55: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 2, 7, 1>, 2, 7>(instr, locals, device); break; case 56: - batch_foreach, 4, 4, 1>, 4, 4>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 57: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 4, 4, 1>, 4, 4>(instr, locals, device); break; case 58: - batch_foreach, 4, 4, 1>, 4, 4>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 59: - batch_foreach, 8, 3, 1>, 8, 3>(instr, locals, device); + batch_foreach, 4, 4, 1>, 4, 4>(instr, locals, device); break; case 60: - batch_foreach, 7, 4, 1>, 7, 4>(instr, locals, device); + batch_foreach, 8, 3, 1>, 8, 3>(instr, locals, device); break; case 61: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 7, 4, 1>, 7, 4>(instr, locals, device); break; case 62: - batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 63: - batch_foreach, 9, 4, 1>, 9, 4>(instr, locals, device); + batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); break; case 64: - batch_foreach, 3, 10, 1>, 3, 10>(instr, locals, device); + batch_foreach, 9, 4, 1>, 9, 4>(instr, locals, device); break; case 65: - batch_foreach, 10, 4, 1>, 10, 4>(instr, locals, device); + batch_foreach, 3, 10, 1>, 3, 10>(instr, locals, device); break; case 66: - batch_foreach, 3, 11, 1>, 3, 11>(instr, locals, device); + batch_foreach, 10, 4, 1>, 10, 4>(instr, locals, device); break; case 67: - batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); + batch_foreach, 3, 11, 1>, 3, 11>(instr, locals, device); break; case 68: - batch_foreach, 4, 3, 1>, 4, 3>(instr, locals, device); + batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); break; case 69: - batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); + batch_foreach, 4, 3, 1>, 4, 3>(instr, locals, device); break; case 70: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); break; case 71: - batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 72: - batch_foreach, 7, 3, 1>, 7, 3>(instr, locals, device); + batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); break; case 73: - batch_foreach, 7, 2, 1>, 7, 2>(instr, locals, device); + batch_foreach, 7, 3, 1>, 7, 3>(instr, locals, device); break; case 74: - batch_foreach, 8, 3, 1>, 8, 3>(instr, locals, device); + batch_foreach, 7, 2, 1>, 7, 2>(instr, locals, device); break; case 75: - batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); + batch_foreach, 8, 3, 1>, 8, 3>(instr, locals, device); break; case 76: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); break; case 77: - batch_foreach, 10, 2, 1>, 10, 2>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 78: - batch_foreach, 10, 3, 1>, 10, 3>(instr, locals, device); + batch_foreach, 10, 2, 1>, 10, 2>(instr, locals, device); break; case 79: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 10, 3, 1>, 10, 3>(instr, locals, device); break; case 80: - batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 81: - batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); + batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); break; case 82: - batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); + batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); break; case 83: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); break; case 84: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 85: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 86: - batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); + batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); break; case 87: - batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); + batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); break; case 88: - batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); + batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); break; case 89: - batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); + batch_foreach, 5, 3, 1>, 5, 3>(instr, locals, device); break; case 90: - batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); + batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); break; case 91: - batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); + batch_foreach, 4, 3, 1>, 4, 3>(instr, locals, device); break; case 92: - batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); + batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); break; case 93: - batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); + batch_foreach, 5, 3, 1>, 5, 3>(instr, locals, device); break; case 94: - batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); + batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); break; case 95: - batch_foreach, 2, 3, 1>, 2, 3>(instr, locals, device); + batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); break; case 96: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 2, 3, 1>, 2, 3>(instr, locals, device); break; case 97: - batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); + batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); break; case 98: - batch_foreach, 3, 3, 1>, 3, 3>(instr, locals, device); + batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); break; case 99: - batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); + batch_foreach, 3, 3, 1>, 3, 3>(instr, locals, device); break; case 100: - batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); + batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); break; case 101: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); break; case 102: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 103: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 104: - batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 105: - batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); break; case 106: - batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); break; case 107: - batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); break; case 108: - batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); + batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); break; case 109: - batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); + batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); break; case 110: - op_matrix_element(instr, locals, device); + batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); break; case 111: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + op_matrix_element(instr, locals, device); break; case 112: - batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 113: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); break; case 114: - op_matmul(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 115: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + op_matmul(instr, locals, device); break; case 116: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 117: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 118: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 119: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 120: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 121: - op_rqs_reshape(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 122: - batch_foreach, 4, 1, 2>, 4, 1>(instr, locals, device); + op_rqs_reshape(instr, locals, device); break; case 123: - batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); + batch_foreach, 4, 1, 2>, 4, 1>(instr, locals, device); break; case 124: - batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); + batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); break; case 125: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); break; case 126: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 127: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 128: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); break; case 129: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); break; case 130: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); break; case 131: - op_discrete_histogram(instr, locals, device); + batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); break; case 132: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + op_discrete_histogram(instr, locals, device); break; case 133: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 134: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 135: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 136: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 137: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 138: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 139: - batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 140: - op_quantile(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 141: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); break; case 142: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + op_quantile(instr, locals, device); break; case 143: - batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 144: - batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 145: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); break; case 146: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); break; case 147: - op_nonzero(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 148: - op_batch_gather(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 149: - op_batch_scatter(instr, locals, device); + op_nonzero(instr, locals, device); break; case 150: - op_random(instr, locals, device); + op_batch_gather(instr, locals, device); break; case 151: - op_random_int(instr, locals, device); + op_batch_scatter(instr, locals, device); break; case 152: - op_unweight(instr, locals, device); + op_random(instr, locals, device); break; case 153: - batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); + op_random_int(instr, locals, device); break; case 154: - batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); + op_unweight(instr, locals, device); break; case 155: - op_vegas_histogram(instr, locals, device); + batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); break; case 156: + batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); + break; +case 157: + op_vegas_histogram(instr, locals, device); + break; +case 158: op_histogram(instr, locals, device); break; diff --git a/madspace/src/kernels/invariants.hpp b/madspace/src/kernels/invariants.hpp index 1710b0111..4eeeae410 100644 --- a/madspace/src/kernels/invariants.hpp +++ b/madspace/src/kernels/invariants.hpp @@ -29,6 +29,7 @@ KERNELSPEC void kernel_breit_wigner_invariant( FIn s_min, FIn s_max, FOut s, + FOut virt, FOut gs ) { auto m2 = mass * mass; @@ -36,11 +37,10 @@ KERNELSPEC void kernel_breit_wigner_invariant( auto y1 = atan((s_min - m2) / gm); auto y2 = atan((s_max - m2) / gm); auto dy21 = y2 - y1; - auto _s = gm * tan(y1 + dy21 * r) + m2; - auto s_sub_m2 = _s - m2; - s = _s; - gs = dy21 * (s_sub_m2 * s_sub_m2 + gm * gm) / gm; + virt = gm * tan(y1 + dy21 * r); + s = virt + m2; + gs = dy21 * (virt * virt + gm * gm) / gm; } template @@ -71,14 +71,17 @@ KERNELSPEC void kernel_stable_invariant( FIn s_min, FIn s_max, FOut s, + FOut virt, FOut gs ) { auto m2 = mass * mass - 1e-2; auto q_max = s_max - m2; auto q_min = s_min - m2; + auto reg_virt = pow(q_max, r) * pow(q_min, 1 - r); - s = pow(q_max, r) * pow(q_min, 1 - r) + m2; - gs = (s - m2) * log(q_max / q_min); + virt = reg_virt - 1e-2; + s = reg_virt + m2; + gs = reg_virt * log(q_max / q_min); } template @@ -108,6 +111,7 @@ KERNELSPEC void kernel_stable_invariant_nu( FIn s_min, FIn s_max, FOut s, + FOut virt, FOut gs ) { auto m2 = mass * mass - 1e-2; @@ -116,10 +120,11 @@ KERNELSPEC void kernel_stable_invariant_nu( auto power = 1.0 - nu; auto qmaxpow = pow(q_max, power); auto qminpow = pow(q_min, power); - auto _s = pow(r * qmaxpow + (1 - r) * qminpow, 1 / power) + m2; + auto reg_virt = pow(r * qmaxpow + (1 - r) * qminpow, 1 / power); - s = _s; - gs = (qmaxpow - qminpow) * pow(_s - m2, nu) / power; + virt = reg_virt - 1e-2; + s = reg_virt + m2; + gs = (qmaxpow - qminpow) * pow(reg_virt, nu) / power; } template diff --git a/madspace/src/kernels/math.hpp b/madspace/src/kernels/math.hpp index b600f2975..1b2f73f91 100644 --- a/madspace/src/kernels/math.hpp +++ b/madspace/src/kernels/math.hpp @@ -49,6 +49,11 @@ KERNELSPEC void backward_kernel_sub( in2_grad += -out_grad; } +template +KERNELSPEC void kernel_neg(FIn in, FOut out) { + out = -in; +} + template KERNELSPEC void kernel_mul(FIn in1, FIn in2, FOut out) { out = in1 * in2; diff --git a/madspace/src/kernels/multichannel.hpp b/madspace/src/kernels/multichannel.hpp index 836eef7f9..bdd301c02 100644 --- a/madspace/src/kernels/multichannel.hpp +++ b/madspace/src/kernels/multichannel.hpp @@ -140,5 +140,22 @@ KERNELSPEC void kernel_permute_momenta( } } +template +KERNELSPEC void kernel_permute_bits( + IIn input, IIn permutations, IIn index, IOut output +) { + auto perm = permutations[single_index(index)]; + for (std::size_t j = 0; j < input.size(); ++j) { + IVal val_in = input[j]; + IVal val_out = val_in; + for (std::size_t i = 0; i < perm.size(); ++i) { + IVal perm_i = perm[i]; + IVal bit_in = (val_in & (1 << perm_i)) >> perm_i; + val_out = (val_out & ~(1 << i)) | (bit_in << i); + } + output[j] = val_out; + } +} + } // namespace kernels } // namespace madspace diff --git a/madspace/src/phasespace/integrand.cpp b/madspace/src/phasespace/integrand.cpp index 2d808ee34..1a644ac58 100644 --- a/madspace/src/phasespace/integrand.cpp +++ b/madspace/src/phasespace/integrand.cpp @@ -2,6 +2,7 @@ #include "madspace/util.hpp" +#include #include using namespace madspace; @@ -207,6 +208,9 @@ NamedVector Integrand::compute_channel_part_ret_types() const { auto acc_float_array = [](int n) { return Type(DataType::dt_float, acc_batch_size, {n}); }; + auto acc_int_array = [](int n) { + return Type(DataType::dt_int, acc_batch_size, {n}); + }; auto acc_four_vec_array = [](int n) { return Type(DataType::dt_float, acc_batch_size, {n, 4}); }; @@ -242,6 +246,12 @@ NamedVector Integrand::compute_channel_part_ret_types() const { } ret.push_back("x1_acc", acc_float); ret.push_back("x2_acc", acc_float); + if (_mapping.return_invariants()) { + int invariant_count = static_cast(_mapping.invariant_count()); + ret.push_back("invariant_pids_and_masks_acc", acc_int_array(invariant_count)); + ret.push_back("invariant_masses_acc", acc_float_array(invariant_count)); + ret.push_back("invariant_virtualities_acc", acc_float_array(invariant_count)); + } ret.push_back("flavor_id", acc_int); ret.push_back("weight_after_cuts", acc_float); if (_madnis_training && !std::holds_alternative(_discrete_after)) { @@ -377,6 +387,16 @@ NamedVector Integrand::build_channel_part( std::array x_acc{ {fb.batch_gather(indices_acc, x0), fb.batch_gather(indices_acc, x1)} }; + Value invariant_pids_and_masks_acc, invariant_masses_acc, + invariant_virtualities_acc; + if (_mapping.return_invariants()) { + invariant_pids_and_masks_acc = + fb.batch_gather(indices_acc, mapping_result["invariant_pids_and_masks"]); + invariant_masses_acc = + fb.batch_gather(indices_acc, mapping_result["invariant_masses"]); + invariant_virtualities_acc = + fb.batch_gather(indices_acc, mapping_result["invariant_virtualities"]); + } for (auto& cond : flow_conditions) { cond = fb.batch_gather(indices_acc, cond); } @@ -529,6 +549,11 @@ NamedVector Integrand::build_channel_part( } out.push_back("x1_acc", x_acc.at(0)); out.push_back("x2_acc", x_acc.at(1)); + if (_mapping.return_invariants()) { + out.push_back("invariant_pids_and_masks_acc", invariant_pids_and_masks_acc); + out.push_back("invariant_masses_acc", invariant_masses_acc); + out.push_back("invariant_virtualities_acc", invariant_virtualities_acc); + } out.push_back("flavor_id", flavor_id); out.push_back("weight_after_cuts", weight_after_cuts); if (_madnis_training && extra_weight_after_cuts) { @@ -597,6 +622,22 @@ NamedVector Integrand::build_common_part( momenta_acc, _flavor_remap.size() > 0 ? fb.gather_int(flavor_id, _flavor_remap) : flavor_id, }; + if (_mapping.return_invariants()) { + // only forward the invariant data if the matrix element actually + // declared inputs for it + auto external_inputs = _diff_xs.matrix_element().external_inputs(); + bool wants_invariants = + std::find( + external_inputs.begin(), + external_inputs.end(), + MatrixElement::invariant_pids_and_masks_in + ) != external_inputs.end(); + if (wants_invariants) { + xs_args.push_back(args.at("invariant_pids_and_masks_acc")); + xs_args.push_back(args.at("invariant_masses_acc")); + xs_args.push_back(args.at("invariant_virtualities_acc")); + } + } xs_args.push_back(x1_acc); xs_args.push_back(x2_acc); xs_args.push_back(flavor_id); diff --git a/madspace/src/phasespace/invariants.cpp b/madspace/src/phasespace/invariants.cpp index 72e701da9..7129d4ca2 100644 --- a/madspace/src/phasespace/invariants.cpp +++ b/madspace/src/phasespace/invariants.cpp @@ -2,16 +2,19 @@ using namespace madspace; -Invariant::Invariant(double power, double mass, double width) : +Invariant::Invariant(double power, double mass, double width, bool return_virtuality) : Mapping( "Invariant", {{"random", batch_float}}, - {{"invariant", batch_float}}, + return_virtuality + ? NamedVector{{"invariant", batch_float}, {"virtuality", batch_float}} + : NamedVector{{"invariant", batch_float}}, {{"invariant_min", batch_float}, {"invariant_max", batch_float}} ), _power(power), _mass(mass), - _width(width) {} + _width(width), + _return_virtuality(return_virtuality) {} Mapping::Result Invariant::build_forward_impl( FunctionBuilder& fb, @@ -19,13 +22,23 @@ Mapping::Result Invariant::build_forward_impl( const NamedVector& conditions ) const { auto r = inputs[0], s_min = conditions[0], s_max = conditions[1]; - auto [s, det] = _width != 0 - ? fb.breit_wigner_invariant(r, _mass, _width, s_min, s_max) - : _power == 0 ? fb.uniform_invariant(r, s_min, s_max) - : _power == 1 - ? fb.stable_invariant(r, _mass, s_min, s_max) - : fb.stable_invariant_nu(r, _mass, _power, s_min, s_max); - return {{{"invariant", s}}, det}; + Value s, virt, det; + if (_width != 0) { + std::tie(s, virt, det) = + fb.breit_wigner_invariant(r, _mass, _width, s_min, s_max); + } else if (_power == 0) { + std::tie(s, det) = fb.uniform_invariant(r, s_min, s_max); + virt = s; + } else if (_power == 1) { + std::tie(s, virt, det) = fb.stable_invariant(r, _mass, s_min, s_max); + } else { + std::tie(s, virt, det) = fb.stable_invariant_nu(r, _mass, _power, s_min, s_max); + } + if (_return_virtuality) { + return {{{"invariant", s}, {"virtuality", virt}}, det}; + } else { + return {{{"invariant", s}}, det}; + } } Mapping::Result Invariant::build_inverse_impl( @@ -34,11 +47,17 @@ Mapping::Result Invariant::build_inverse_impl( const NamedVector& conditions ) const { auto s = inputs[0], s_min = conditions[0], s_max = conditions[1]; - auto [r, det] = _width != 0 - ? fb.breit_wigner_invariant_inverse(s, _mass, _width, s_min, s_max) - : _power == 0 ? fb.uniform_invariant_inverse(s, s_min, s_max) - : _power == 1 - ? fb.stable_invariant_inverse(s, _mass, s_min, s_max) - : fb.stable_invariant_nu_inverse(s, _mass, _power, s_min, s_max); + Value r, det; + if (_width != 0) { + std::tie(r, det) = + fb.breit_wigner_invariant_inverse(s, _mass, _width, s_min, s_max); + } else if (_power == 0) { + std::tie(r, det) = fb.uniform_invariant_inverse(s, s_min, s_max); + } else if (_power == 1) { + std::tie(r, det) = fb.stable_invariant_inverse(s, _mass, s_min, s_max); + } else { + std::tie(r, det) = + fb.stable_invariant_nu_inverse(s, _mass, _power, s_min, s_max); + } return {{{"random", r}}, det}; } diff --git a/madspace/src/phasespace/matrix_element.cpp b/madspace/src/phasespace/matrix_element.cpp index 178910150..35cfb8352 100644 --- a/madspace/src/phasespace/matrix_element.cpp +++ b/madspace/src/phasespace/matrix_element.cpp @@ -8,7 +8,8 @@ MatrixElement::MatrixElement( const std::vector& inputs, const std::vector& outputs, std::size_t diagram_count, - bool sample_random_inputs + bool sample_random_inputs, + std::size_t invariant_count ) : FunctionGenerator( "MatrixElement", @@ -53,6 +54,24 @@ MatrixElement::MatrixElement( case channel_in: arg_types.push_back("channel", batch_int); break; + case invariant_count_in: + // host constant, not a graph value; see build_function_impl + break; + case invariant_pids_and_masks_in: + arg_types.push_back( + "invariant_pids_and_masks", batch_int_array(invariant_count) + ); + break; + case invariant_masses_in: + arg_types.push_back( + "invariant_masses", batch_float_array(invariant_count) + ); + break; + case invariant_virtualities_in: + arg_types.push_back( + "invariant_virtualities", batch_float_array(invariant_count) + ); + break; default: throw std::invalid_argument("unknown input type"); } @@ -148,11 +167,25 @@ NamedVector MatrixElement::build_function_impl( case channel_in: input_key = UMAMI_IN_CHANNEL_INDEX; break; + case invariant_count_in: + input_key = UMAMI_IN_INVARIANT_COUNT; + break; + case invariant_pids_and_masks_in: + input_key = UMAMI_IN_INVARIANT_PIDS_AND_MASKS; + break; + case invariant_masses_in: + input_key = UMAMI_IN_INVARIANT_MASSES; + break; + case invariant_virtualities_in: + input_key = UMAMI_IN_INVARIANT_VIRTUALITIES; + break; } matrix_args.push_back(static_cast(input_key)); - if (_sample_random_inputs && - (input == random_color_in || input == random_helicity_in || - input == random_diagram_in)) { + if (input == invariant_count_in) { + matrix_args.push_back(static_cast(invariant_count())); + } else if (_sample_random_inputs && + (input == random_color_in || input == random_helicity_in || + input == random_diagram_in)) { matrix_args.push_back(random.at(random_index)); ++random_index; } else { @@ -191,6 +224,11 @@ NamedVector MatrixElement::build_function_impl( std::vector MatrixElement::external_inputs() const { std::vector ret; for (auto input : _inputs) { + // invariant_count_in is a host constant (see build_function_impl), not + // supplied as an external graph value + if (input == invariant_count_in) { + continue; + } if (!_sample_random_inputs || (input != random_color_in && input != random_helicity_in && input != random_diagram_in)) { diff --git a/madspace/src/phasespace/phasespace.cpp b/madspace/src/phasespace/phasespace.cpp index f95f2e165..9a50520e5 100644 --- a/madspace/src/phasespace/phasespace.cpp +++ b/madspace/src/phasespace/phasespace.cpp @@ -80,9 +80,6 @@ nested_vector2 invert_permutations(nested_vector2 perms_in) return perms_out; } -} // namespace - -namespace { // Chain (color) order for the t-channel ColorOrderedMapping: the externally // supplied order if given, else the default single chain [0, 2, ..., n+1, 1]. std::vector ps_chain_order( @@ -116,6 +113,29 @@ std::size_t ps_discrete_dim( } return 0; } + +// Total number of invariants reported when return_invariants is set. +std::size_t ps_invariant_count( + const Topology& topology, + bool leptonic, + PhaseSpaceMapping::TChannelMode t_channel_mode +) { + std::size_t invariant_count = + topology.decays().size() - topology.outgoing_masses().size(); + if (leptonic || + (topology.t_propagator_count() != 0 && + t_channel_mode == PhaseSpaceMapping::chili)) { + --invariant_count; + } + if (topology.t_propagator_count() > 0 && + (t_channel_mode == PhaseSpaceMapping::propagator || + topology.t_propagator_count() < 2)) { + invariant_count += + TPropagatorMapping::invariant_count(topology.t_propagator_count()); + } + return invariant_count; +} + } // namespace PhaseSpaceMapping::PhaseSpaceMapping( @@ -126,7 +146,8 @@ PhaseSpaceMapping::PhaseSpaceMapping( TChannelMode t_channel_mode, const std::optional& cuts, const std::vector>& permutations, - const std::optional>& color_order + const std::optional>& color_order, + bool return_invariants ) : Mapping( "PhaseSpaceMapping", @@ -148,9 +169,26 @@ PhaseSpaceMapping::PhaseSpaceMapping( } return in; }(), - {{"momenta", batch_four_vec_array(topology.outgoing_masses().size() + 2)}, - {"x1", batch_float}, - {"x2", batch_float}}, + [&] { + NamedVector out{ + {"momenta", + batch_four_vec_array(topology.outgoing_masses().size() + 2)}, + {"x1", batch_float}, + {"x2", batch_float}, + }; + if (return_invariants) { + std::size_t invariant_count = + ps_invariant_count(topology, leptonic, t_channel_mode); + out.push_back( + "invariant_pids_and_masks", batch_int_array(invariant_count) + ); + out.push_back("invariant_masses", batch_float_array(invariant_count)); + out.push_back( + "invariant_virtualities", batch_float_array(invariant_count) + ); + } + return out; + }(), permutations.size() > 1 ? NamedVector{{"permutation_index", batch_int}} : NamedVector{} @@ -167,7 +205,8 @@ PhaseSpaceMapping::PhaseSpaceMapping( (_topology.t_propagator_count() == 0 || t_channel_mode != PhaseSpaceMapping::chili) ), - _t_mapping(std::monostate{}) { + _t_mapping(std::monostate{}), + _return_invariants(return_invariants) { bool has_t_channel = _topology.t_propagator_count() > 0; struct DecayInfo { double m_min, pt_min, eta_max; @@ -212,7 +251,8 @@ PhaseSpaceMapping::PhaseSpaceMapping( if (!is_com_decay || _map_luminosity) { double mass = decay.width == 0. ? 0. : decay.mass; double width = decay.width; - info.invariant = Invariant(invariant_power, mass, width); + info.invariant = + Invariant(invariant_power, mass, width, _return_invariants); } } for (std::size_t index : _topology.decay_integration_order()) { @@ -299,7 +339,10 @@ PhaseSpaceMapping::PhaseSpaceMapping( } else if (t_channel_mode == PhaseSpaceMapping::propagator || topology.t_propagator_count() < 2) { _t_mapping = TPropagatorMapping( - _topology.t_integration_order(), invariant_power, pt_min + _topology.t_integration_order(), + invariant_power, + pt_min, + _return_invariants ); } else if (t_channel_mode == PhaseSpaceMapping::rambo) { // TODO: add massless special case @@ -322,7 +365,8 @@ PhaseSpaceMapping::PhaseSpaceMapping( double invariant_power, TChannelMode mode, const std::optional& cuts, - const std::optional>& color_order + const std::optional>& color_order, + bool return_invariants ) : PhaseSpaceMapping( Topology([&] { @@ -361,7 +405,8 @@ PhaseSpaceMapping::PhaseSpaceMapping( mode, cuts, {}, - color_order + color_order, + return_invariants ) {} Mapping::Result PhaseSpaceMapping::build_forward_impl( @@ -384,6 +429,7 @@ Mapping::Result PhaseSpaceMapping::build_forward_impl( ValueVec dets{_pi_factors}; Value x1 = 1.0, x2 = 1.0; + ValueVec invariant_pids_and_masks, invariant_masses, invariant_virtualities; // initialize masses and square masses std::vector decay_data( @@ -414,6 +460,14 @@ Mapping::Result PhaseSpaceMapping::build_forward_impl( auto invariant = _s_invariants.at(invariant_index++) .build_forward(fb, {next_random()}, {s_min, s_max}); + + if (_return_invariants) { + me_int_t pid = decay.mass == 0 ? 0 : decay.pdg_id; + invariant_pids_and_masks.push_back((pid << 16) + decay.momentum_mask); + invariant_masses.push_back(invariant["invariant"]); + invariant_virtualities.push_back(invariant["virtuality"]); + } + data.mass2 = invariant["invariant"]; data.mass = fb.sqrt(data.mass2.value()); dets.push_back(invariant["det"]); @@ -476,6 +530,25 @@ Mapping::Result PhaseSpaceMapping::build_forward_impl( x1 = x1_new; x2 = x2_new; } + if constexpr (std::is_same_v) { + if (_return_invariants) { + std::vector leg_masks; + for (std::size_t child_index : + _topology.decays().at(0).child_indices) { + leg_masks.push_back( + _topology.decays().at(child_index).momentum_mask + ); + } + auto t_masks = t_mapping.invariant_masks(leg_masks); + for (std::size_t j = 0; j < t_masks.size(); ++j) { + Value t_invariant = t_result.at(result_index + j); + // all invariants in TPropagatorMapping are massless + invariant_pids_and_masks.push_back(t_masks.at(j)); + invariant_masses.push_back(t_invariant); + invariant_virtualities.push_back(t_invariant); + } + } + } }, [&](std::monostate) { auto [p1, p2] = fb.com_p_in(sqrt_s_hat); @@ -522,12 +595,11 @@ Mapping::Result PhaseSpaceMapping::build_forward_impl( auto p_ext_stack = fb.stack(p_ext); // permute momenta if permutations are given + bool has_unsorted_perm = _permutations.size() == 1 && + !std::is_sorted(_permutations.at(0).begin(), _permutations.at(0).end()); if (_permutations.size() > 1) { p_ext_stack = fb.permute_momenta(p_ext_stack, _permutations, conditions.at(0)); - } else if (_permutations.size() == 1 && - !std::is_sorted( - _permutations.at(0).begin(), _permutations.at(0).end() - )) { + } else if (has_unsorted_perm) { p_ext_stack = fb.permute_momenta(p_ext_stack, _permutations, static_cast(0)); } @@ -536,7 +608,30 @@ Mapping::Result PhaseSpaceMapping::build_forward_impl( auto p_ext_lab = _map_luminosity ? fb.boost_beam(p_ext_stack, x1, x2) : p_ext_stack; dets.push_back(_cuts.build_function(fb, {p_ext_lab}).at(0)); auto ps_weight = fb.cut_unphysical(fb.product(dets), p_ext_lab, x1, x2); - return {{{"momenta", p_ext_lab}, {"x1", x1}, {"x2", x2}}, ps_weight}; + + if (_return_invariants) { + Value pids_and_masks = fb.stack(invariant_pids_and_masks); + if (_permutations.size() > 1) { + pids_and_masks = + fb.permute_bits(pids_and_masks, _permutations, conditions.at(0)); + } else if (has_unsorted_perm) { + pids_and_masks = fb.permute_bits( + pids_and_masks, _permutations, static_cast(0) + ); + } + + return { + {{"momenta", p_ext_lab}, + {"x1", x1}, + {"x2", x2}, + {"invariant_pids_and_masks", pids_and_masks}, + {"invariant_masses", fb.stack(invariant_masses)}, + {"invariant_virtualities", fb.stack(invariant_virtualities)}}, + ps_weight + }; + } else { + return {{{"momenta", p_ext_lab}, {"x1", x1}, {"x2", x2}}, ps_weight}; + } } Mapping::Result PhaseSpaceMapping::build_inverse_impl( diff --git a/madspace/src/phasespace/t_propagator_mapping.cpp b/madspace/src/phasespace/t_propagator_mapping.cpp index af1dc70b2..3a427a7cb 100644 --- a/madspace/src/phasespace/t_propagator_mapping.cpp +++ b/madspace/src/phasespace/t_propagator_mapping.cpp @@ -22,7 +22,8 @@ static bool has_pt_cut(const std::vector& pt_min) { TPropagatorMapping::TPropagatorMapping( const std::vector& integration_order, double invariant_power, - const std::vector& pt_min + const std::vector& pt_min, + bool return_invariants ) : Mapping( "TPropagatorMapping", @@ -38,6 +39,14 @@ TPropagatorMapping::TPropagatorMapping( for (std::size_t i = 0; i < integration_order.size() + 3; ++i) { output_types.push_back(std::format("momentum{}", i), batch_four_vec); } + if (return_invariants) { + // one invariant per t-propagator, followed by one per + // intermediate recoil-system invariant (see invariant_masks) + for (std::size_t i = 0; i < invariant_count(integration_order.size()); + ++i) { + output_types.push_back(std::format("invariant{}", i), batch_float); + } + } return output_types; }(), [&] { @@ -51,8 +60,13 @@ TPropagatorMapping::TPropagatorMapping( _integration_order(integration_order), _pt_min(pt_min), _has_cut(has_pt_cut(pt_min)), - _com_scattering(true, invariant_power, 0., 0., has_pt_cut(pt_min)), - _lab_scattering(false, invariant_power, 0., 0., has_pt_cut(pt_min)) { + _return_invariants(return_invariants), + _com_scattering( + true, invariant_power, 0., 0., has_pt_cut(pt_min), return_invariants + ), + _lab_scattering( + false, invariant_power, 0., 0., has_pt_cut(pt_min), return_invariants + ) { std::size_t next_index_low = 0; std::size_t next_index_high = integration_order.size() - 1; for (std::size_t index : integration_order) { @@ -68,6 +82,38 @@ TPropagatorMapping::TPropagatorMapping( } } +std::vector +TPropagatorMapping::invariant_masks(const std::vector& outgoing_masks) const { + std::vector masks; + + // one mask per t-channel propagator, in physical order along the chain + // (see t_invariants.at(index) below) + int mask = 1; + for (std::size_t i = 0; i < _integration_order.size(); ++i) { + mask |= outgoing_masks.at(i); + masks.push_back(mask); + } + + // one mask per intermediate recoil-system invariant, mirroring the + // mass_sum_invariants computation below + if (_integration_order.size() > 1) { + std::size_t last_index = _integration_order.back(); + std::vector min_masks{ + outgoing_masks.at(last_index) | outgoing_masks.at(last_index + 1) + }; + for (std::size_t i = _integration_order.size() - 2; i > 0; --i) { + int next_mask = + outgoing_masks.at(_integration_order.at(i) + _sample_sides.at(i)); + min_masks.push_back(min_masks.back() | next_mask); + } + for (int recoil_mask : std::views::reverse(min_masks)) { + masks.push_back(recoil_mask); + } + } + + return masks; +} + Mapping::Result TPropagatorMapping::build_forward_impl( FunctionBuilder& fb, const NamedVector& inputs, @@ -80,6 +126,7 @@ Mapping::Result TPropagatorMapping::build_forward_impl( ValueVec dets; ValueVec mass_sum_invariants; + ValueVec s_invariants; if (_integration_order.size() > 1) { // compute sums of outgoing masses, starting from those sampled last std::size_t last_index = _integration_order.back(); @@ -106,6 +153,9 @@ Mapping::Result TPropagatorMapping::build_forward_impl( auto mass = fb.sqrt(s_result["invariant"]); mass_sum_invariants.push_back(mass); dets.push_back(s_result["det"]); + if (_return_invariants) { + s_invariants.push_back(s_result["invariant"]); + } max_mass = mass; } } @@ -130,6 +180,10 @@ Mapping::Result TPropagatorMapping::build_forward_impl( } // sample t-invariants and build momenta of t-channel part of the diagram + ValueVec t_invariants; + if (_return_invariants) { + t_invariants.resize(_integration_order.size()); + } Value k_rest; bool first = true; for (auto [index, side, mass_sum] : @@ -152,6 +206,9 @@ Mapping::Result TPropagatorMapping::build_forward_impl( auto k = ks.at(1); p_ext.at(sampled_index + 2) = k; dets.push_back(ks["det"]); + if (_return_invariants) { + t_invariants.at(index) = fb.neg(ks["invariant"]); + } if (side) { p2_rest = fb.sub(p2_rest, k); } else { @@ -159,7 +216,11 @@ Mapping::Result TPropagatorMapping::build_forward_impl( } } p_ext.at(_integration_order.back() + 2) = k_rest; - return {{output_types().keys(), p_ext}, fb.product(dets)}; + + ValueVec outputs = p_ext; + outputs.insert(outputs.end(), t_invariants.begin(), t_invariants.end()); + outputs.insert(outputs.end(), s_invariants.begin(), s_invariants.end()); + return {{output_types().keys(), outputs}, fb.product(dets)}; } Mapping::Result TPropagatorMapping::build_inverse_impl( diff --git a/madspace/src/phasespace/topology.cpp b/madspace/src/phasespace/topology.cpp index 90fbac951..988c04d0d 100644 --- a/madspace/src/phasespace/topology.cpp +++ b/madspace/src/phasespace/topology.cpp @@ -90,28 +90,29 @@ void build_decays( std::size_t decay_index = decays.size(); switch (line_ref.type()) { case Diagram::outgoing: - decays.push_back( - {decay_index, - parent_decay_index, - {}, - diagram.outgoing_masses().at(line_ref.index()), - 0., - 0., - 0.} - ); + decays.push_back({ + .index = decay_index, + .parent_index = parent_decay_index, + .child_indices = {}, + .mass = diagram.outgoing_masses().at(line_ref.index()), + .width = 0., + .e_min = 0., + .e_max = 0., + .pdg_id = 0, + }); outgoing_indices.at(line_ref.index()) = decay_index; break; case Diagram::propagator: { auto& propagator = diagram.propagators().at(line_ref.index()); decays.push_back({ - decay_index, - parent_decay_index, - {}, - propagator.mass, - propagator.width, - propagator.e_min, - propagator.e_max, - propagator.pdg_id, + .index = decay_index, + .parent_index = parent_decay_index, + .child_indices = {}, + .mass = propagator.mass, + .width = propagator.width, + .e_min = propagator.e_min, + .e_max = propagator.e_max, + .pdg_id = propagator.pdg_id, }); decay_indices.push_back(decay_index); integration_order.push_back(propagator.integration_order); @@ -143,6 +144,26 @@ void build_decays( } } +void build_decay_momentum_masks( + const std::vector& outgoing_indices, + std::vector& decays +) { + for (std::size_t ext_index = 2; std::size_t index : outgoing_indices) { + decays.at(index).momentum_mask = 1 << ext_index; + ++ext_index; + } + for (auto& decay : std::views::reverse(decays)) { + if (decay.child_indices.empty()) { + continue; + } + int momentum_mask = 0; + for (std::size_t child_index : decay.child_indices) { + momentum_mask |= decays.at(child_index).momentum_mask; + } + decay.momentum_mask = momentum_mask; + } +} + std::string decay_label( const Topology::Decay& decay, const std::unordered_map& decay_order, @@ -359,6 +380,8 @@ std::vector Topology::topologies(const Diagram& diagram) { } } + build_decay_momentum_masks(topo._outgoing_indices, topo._decays); + std::size_t massive_decays = 0; for (std::size_t index : decay_indices) { if (topo._decays.at(index).mass != 0) { diff --git a/madspace/src/phasespace/two_particle.cpp b/madspace/src/phasespace/two_particle.cpp index 4e118b66e..dfc7d1342 100644 --- a/madspace/src/phasespace/two_particle.cpp +++ b/madspace/src/phasespace/two_particle.cpp @@ -69,7 +69,12 @@ Mapping::Result TwoBodyDecay::build_inverse_impl( } TwoToTwoParticleScattering::TwoToTwoParticleScattering( - bool com, double invariant_power, double mass, double width, bool has_cut + bool com, + double invariant_power, + double mass, + double width, + bool has_cut, + bool return_invariant ) : Mapping( "TwoToTwoParticleScattering", @@ -77,7 +82,15 @@ TwoToTwoParticleScattering::TwoToTwoParticleScattering( {"random_inv", batch_float}, {"mass1", batch_float}, {"mass2", batch_float}}, - {{"momentum1", batch_four_vec}, {"momentum2", batch_four_vec}}, + [&] { + NamedVector out{ + {"momentum1", batch_four_vec}, {"momentum2", batch_four_vec} + }; + if (return_invariant) { + out.push_back("invariant", batch_float); + } + return out; + }(), [&] { NamedVector cond{ {"momentum_in1", batch_four_vec}, {"momentum_in2", batch_four_vec} @@ -91,7 +104,8 @@ TwoToTwoParticleScattering::TwoToTwoParticleScattering( ), _com(com), _invariant(invariant_power, mass, width), - _has_cut(has_cut) {} + _has_cut(has_cut), + _return_invariant(return_invariant) {} Mapping::Result TwoToTwoParticleScattering::build_forward_impl( FunctionBuilder& fb, @@ -112,9 +126,11 @@ Mapping::Result TwoToTwoParticleScattering::build_forward_impl( : fb.two_to_two_particle_scattering( r_phi, p_in1, p_in2, t_result["invariant"], m1, m2 ); - return { - {{"momentum1", p1}, {"momentum2", p2}}, fb.mul(t_result["det"], det_scatter) - }; + NamedVector out{{"momentum1", p1}, {"momentum2", p2}}; + if (_return_invariant) { + out.push_back("invariant", t_result["invariant"]); + } + return {out, fb.mul(t_result["det"], det_scatter)}; } Mapping::Result TwoToTwoParticleScattering::build_inverse_impl( diff --git a/madspace/src/python/instruction_set.hpp b/madspace/src/python/instruction_set.hpp index 73fceb63d..b29c83728 100644 --- a/madspace/src/python/instruction_set.hpp +++ b/madspace/src/python/instruction_set.hpp @@ -30,6 +30,7 @@ void add_instructions(py::classh& fb) { fb.def("add", &FunctionBuilder::add, py::arg("in1"), py::arg("in2")); fb.def("add_int", &FunctionBuilder::add_int, py::arg("in1"), py::arg("in2")); fb.def("sub", &FunctionBuilder::sub, py::arg("in1"), py::arg("in2")); + fb.def("neg", &FunctionBuilder::neg, py::arg("in")); fb.def("mul", &FunctionBuilder::mul, py::arg("in1"), py::arg("in2")); fb.def("div", &FunctionBuilder::div, py::arg("in1"), py::arg("in2")); fb.def("reduce_sum", &FunctionBuilder::reduce_sum, py::arg("in1")); @@ -146,6 +147,7 @@ void add_instructions(py::classh& fb) { fb.def("sample_discrete_probs_inverse", &FunctionBuilder::sample_discrete_probs_inverse, py::arg("index"), py::arg("probs")); fb.def("discrete_histogram", &FunctionBuilder::discrete_histogram, py::arg("input"), py::arg("weights"), py::arg("option_count")); fb.def("permute_momenta", &FunctionBuilder::permute_momenta, py::arg("momenta"), py::arg("permutations"), py::arg("index")); + fb.def("permute_bits", &FunctionBuilder::permute_bits, py::arg("input"), py::arg("permutations"), py::arg("index")); fb.def("gather", &FunctionBuilder::gather, py::arg("index"), py::arg("choices")); fb.def("gather_int", &FunctionBuilder::gather_int, py::arg("index"), py::arg("choices")); fb.def("gather_vector", &FunctionBuilder::gather_vector, py::arg("index"), py::arg("choices")); diff --git a/madspace/src/python/madspace.cpp b/madspace/src/python/madspace.cpp index 5e3ccf0e2..7d58c0ca7 100644 --- a/madspace/src/python/madspace.cpp +++ b/madspace/src/python/madspace.cpp @@ -410,10 +410,11 @@ PYBIND11_MODULE(_madspace_py, m) { py::classh(m, "Invariant") .def( - py::init(), + py::init(), py::arg("power") = 0., py::arg("mass") = 0., - py::arg("width") = 0. + py::arg("width") = 0., + py::arg("return_virtuality") = false ); py::classh(m, "Luminosity") @@ -715,7 +716,8 @@ PYBIND11_MODULE(_madspace_py, m) { PhaseSpaceMapping::TChannelMode, const std::optional&, const nested_vector2&, - const std::optional>&>(), + const std::optional>&, + bool>(), py::arg("topology"), py::arg("cm_energy"), py::arg("leptonic") = false, @@ -723,7 +725,8 @@ PYBIND11_MODULE(_madspace_py, m) { py::arg("t_channel_mode") = PhaseSpaceMapping::propagator, py::arg("cuts") = std::nullopt, py::arg("permutations") = std::vector{}, - py::arg("color_order") = std::nullopt + py::arg("color_order") = std::nullopt, + py::arg("return_invariants") = false ) .def( py::init< @@ -733,19 +736,23 @@ PYBIND11_MODULE(_madspace_py, m) { double, PhaseSpaceMapping::TChannelMode, std::optional, - const std::optional>&>(), + const std::optional>&, + bool>(), py::arg("masses"), py::arg("cm_energy"), py::arg("leptonic") = false, py::arg("invariant_power") = 0.8, py::arg("mode") = PhaseSpaceMapping::rambo, py::arg("cuts") = std::nullopt, - py::arg("color_order") = std::nullopt + py::arg("color_order") = std::nullopt, + py::arg("return_invariants") = false ) .def("random_dim", &PhaseSpaceMapping::random_dim) .def("discrete_dim", &PhaseSpaceMapping::discrete_dim) .def("particle_count", &PhaseSpaceMapping::particle_count) - .def("channel_count", &PhaseSpaceMapping::channel_count); + .def("channel_count", &PhaseSpaceMapping::channel_count) + .def("return_invariants", &PhaseSpaceMapping::return_invariants) + .def("invariant_count", &PhaseSpaceMapping::invariant_count); py::classh(m, "MultiChannelFunction") .def( @@ -767,6 +774,10 @@ PYBIND11_MODULE(_madspace_py, m) { {"helicity_in", MatrixElement::helicity_in}, {"diagram_in", MatrixElement::diagram_in}, {"channel_in", MatrixElement::channel_in}, + {"invariant_count_in", MatrixElement::invariant_count_in}, + {"invariant_pids_and_masks_in", MatrixElement::invariant_pids_and_masks_in}, + {"invariant_masses_in", MatrixElement::invariant_masses_in}, + {"invariant_virtualities_in", MatrixElement::invariant_virtualities_in}, } ); add_enum( @@ -788,28 +799,33 @@ PYBIND11_MODULE(_madspace_py, m) { const std::vector&, const std::vector&, std::size_t, - bool>(), + bool, + std::size_t>(), py::arg("matrix_element_index"), py::arg("particle_count"), py::arg("inputs"), py::arg("outputs"), py::arg("diagram_count") = 1, - py::arg("sample_random_inputs") = false + py::arg("sample_random_inputs") = false, + py::arg("invariant_count") = 0 ) .def( py::init< const MatrixElementApi&, const std::vector&, const std::vector&, - bool>(), + bool, + std::size_t>(), py::arg("matrix_element_api"), py::arg("inputs"), py::arg("outputs"), - py::arg("sample_random_inputs") = false + py::arg("sample_random_inputs") = false, + py::arg("invariant_count") = 0 ) .def("matrix_element_index", &MatrixElement::diagram_count) .def("diagram_count", &MatrixElement::diagram_count) - .def("particle_count", &MatrixElement::particle_count); + .def("particle_count", &MatrixElement::particle_count) + .def("invariant_count", &MatrixElement::invariant_count); py::classh mlp(m, "MLP"); add_enum( diff --git a/madspace/tests/test_processes.py b/madspace/tests/test_processes.py index 70262391e..30f6242c2 100644 --- a/madspace/tests/test_processes.py +++ b/madspace/tests/test_processes.py @@ -11,6 +11,7 @@ BATCH_SIZE = 1000 CM_ENERGY = 13000.0 +INVARIANT_POWER = 0.2 rng = np.random.default_rng(1234) @@ -45,7 +46,10 @@ def mapping(process): ) topology = ms.Topology(diagram) return ms.PhaseSpaceMapping( - topology, CM_ENERGY, permutations=process["permutations"] + topology, + CM_ENERGY, + permutations=process["permutations"], + invariant_power=INVARIANT_POWER, ) @@ -59,6 +63,33 @@ def permutation_count(process): return len(process["permutations"]) +@pytest.fixture +def mapping_with_propagators(process): + # assign a distinct nonzero pdg_id to each distinct nonzero propagator + # mass, so returned propagators can be mapped back to a mass + mass_by_pid = {} + propagators = [] + for mass, width in process["propagators"]: + pid = 0 + if mass != 0.0: + pid = next((p for p, m in mass_by_pid.items() if m == mass), None) + if pid is None: + pid = len(mass_by_pid) + 1 + mass_by_pid[pid] = mass + propagators.append(ms.Propagator(mass, width, 0, 0.0, 0.0, pid)) + diagram = ms.Diagram( + process["incoming_masses"], + process["outgoing_masses"], + propagators, + process["vertices"], + ) + topology = ms.Topology(diagram) + mapping = ms.PhaseSpaceMapping( + topology, CM_ENERGY, return_invariants=True, invariant_power=INVARIANT_POWER + ) + return mapping, mass_by_pid + + def test_process_masses(mapping, masses, permutation_count): r = rng.random((BATCH_SIZE, mapping.random_dim())) condition = ( @@ -121,3 +152,44 @@ def test_process_inverse(mapping, masses, permutation_count): r_inv, det_inv = mapping.map_inverse(map_out, condition) assert r_inv == approx(r, abs=1e-3, rel=1e-3) assert det_inv == approx(1 / det, rel=1e-3) + + +def test_process_propagators(mapping_with_propagators): + mapping, mass_by_pid = mapping_with_propagators + r = rng.random((BATCH_SIZE, mapping.random_dim())) + result = mapping.map_forward([r]) + + p_ext = result.momenta + n_ext = p_ext.shape[1] + # momentum_mask bit i always selects p_ext[:, i]; incoming particles (bits + # 0 and 1) enter with a flipped sign, since a propagator's momentum is the + # incoming momentum minus whatever outgoing momenta have already branched + # off. Squaring removes the resulting overall sign ambiguity. + signs = np.where(np.arange(n_ext) < 2, -1.0, 1.0) + + # pid and momentum mask are diagram-level constants, broadcast over the batch + pids_and_masks = result.invariant_pids_and_masks[0] + invariants = result.invariant_masses + virtualities = result.invariant_virtualities + + for i, pid_and_mask in enumerate(pids_and_masks): + pid = int(pid_and_mask) >> 16 + momentum_mask = int(pid_and_mask) & 0xFFFF + + bits = ((momentum_mask >> np.arange(n_ext)) & 1).astype(bool) + p_sum = np.where(bits[None, :, None], signs[None, :, None] * p_ext, 0.0).sum( + axis=1 + ) + + invariant_from_momenta = p_sum[:, 0] ** 2 - np.sum(p_sum[:, 1:] ** 2, axis=1) + assert invariant_from_momenta == approx(invariants[:, i], rel=1e-5, abs=1e-5) + + if pid == 0: + # pid == 0 stands for a massless particle (or an unused slot), so + # mass == 0 and virtuality == invariant + assert invariants[:, i] == approx(virtualities[:, i], rel=1e-5, abs=1e-5) + else: + virtuality_from_mass = invariant_from_momenta - mass_by_pid[pid] ** 2 + assert virtuality_from_mass == approx( + virtualities[:, i], rel=1e-5, abs=1e-5 + )