diff --git a/madgraph/iolibs/export_cpp.py b/madgraph/iolibs/export_cpp.py index f24eb098c..c974430c9 100755 --- a/madgraph/iolibs/export_cpp.py +++ b/madgraph/iolibs/export_cpp.py @@ -27,6 +27,7 @@ import shutil import subprocess import json +from collections import defaultdict import madgraph.core.base_objects as base_objects import madgraph.core.color_algebra as color @@ -1544,9 +1545,7 @@ def get_sigmaKin_lines(self, color_amplitudes, write=True): return replace_dict def get_flavor_table(self, matrix_element): - print(list(matrix_element.get_external_flavors())) flavors = list(matrix_element.get_external_flavors_with_iden()) - print(flavors) flavor_dict = { 1: 0, 2: 1, 3: 2, 4: 3, # quarks 11: 0, 13: 1, 15: 2, # charged leptons @@ -3199,6 +3198,7 @@ def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.me_lib_format = args[1].get("me_lib_format", None) self.process_info = [] + self.merged_subprocesses = defaultdict(list) def generate_subprocess_directory( self, matrix_element, cpp_helas_call_writer, proc_number=None @@ -3234,7 +3234,11 @@ def generate_subprocess_directory( plot.draw() me_lib_path = self.me_lib_format.format(process_id = proc_dir_name) - self.process_info.append(process_exporter_mg7.get_subprocess_info(dirpath, me_lib_path)) + subproc_info, diagram_tags, subproc_class = process_exporter_mg7.get_subprocess_info(dirpath, me_lib_path) + self.merged_subprocesses[subproc_class].append( + (len(self.process_info), diagram_tags) + ) + self.process_info.append(subproc_info) def copy_template(self, model): super().copy_template(model) @@ -3258,6 +3262,83 @@ def copy_template(self, model): ) os.chmod(madnis_bin, 0o755) + def get_merged_info(self): + merged_subproc_info = [] + for subprocesses in self.merged_subprocesses.values(): + channels = [] + flavors = [] + subproc_indices = [] + unique_diagram_tags = [] + unique_diagrams = [] + for subproc_index_in_group, (subproc_index, diagram_tags) in enumerate( + subprocesses + ): + subproc_indices.append(subproc_index) + subproc_info = self.process_info[subproc_index] + flavor_offset = len(flavors) + flavors.extend( + { + "subprocess": subproc_index_in_group, + "flavor": ps_flavor, + } + for ps_flavor, flavor in enumerate( + subproc_info["flavors"] + ) + ) + + for chan_info, chan_tags in zip(subproc_info["channels"], diagram_tags): + chan_index = len(channels) + + same_diags = [] + for tag in chan_tags: + try: + index = unique_diagram_tags.index(tag) + chan_index, diag_index = unique_diagrams[index] + same_diags.append(diag_index) + except ValueError: + same_diags.append(None) + + if chan_index == len(channels): + channels.append( + { + "subprocess": subproc_index, + "channel": chan_index, + "diagrams": [], + } + ) + + diagrams = channels[chan_index]["diagrams"] + for diag_info, same_diag, tag in zip( + chan_info["diagrams"], same_diags, chan_tags + ): + if same_diag is None: + unique_diagram_tags.append(tag) + diag_index = len(diagrams) + unique_diagrams.append((chan_index, diag_index)) + diagrams.append( + { + "diagram": [-1] * len(subprocesses), + "permutation": diag_info["permutation"], + "active_flavors": [], + } + ) + else: + diag_index = same_diag + diag_dict = diagrams[diag_index] + diag_dict["diagram"][subproc_index_in_group] = diag_info["diagram"] + diag_dict["active_flavors"].extend( + flavor_offset + flav for flav in diag_info["active_flavors"] + ) + + merged_subproc_info.append({ + "incoming": subproc_info["incoming"], + "outgoing": subproc_info["outgoing"], + "subprocesses": subproc_indices, + "channels": channels, + "flavors": flavors, + }) + return merged_subproc_info + def finalize(self, matrix_elements=None, history='', *args, **kwargs): file_name = os.path.normpath(os.path.join( self.dir_path, "SubProcesses", "subprocesses.json" @@ -3265,6 +3346,12 @@ def finalize(self, matrix_elements=None, history='', *args, **kwargs): with open(file_name, 'w') as f: json.dump(self.process_info, f) + merged_file_name = os.path.normpath(os.path.join( + self.dir_path, "SubProcesses", "merged_subprocesses.json" + )) + with open(merged_file_name, 'w') as f: + json.dump(self.get_merged_info(), f) + # Generate Cards/run_card.toml from the template, filling in # process-dependent defaults (mirrors the LO run_card.dat logic). self.create_run_card(matrix_elements, history) diff --git a/madgraph/iolibs/export_mg7.py b/madgraph/iolibs/export_mg7.py index d2cd624d6..d505c9061 100644 --- a/madgraph/iolibs/export_mg7.py +++ b/madgraph/iolibs/export_mg7.py @@ -4,6 +4,50 @@ from madgraph.various.diagram_symmetry import find_symmetry, IdentifySGConfigTag from madgraph.iolibs import export_cpp +from madgraph.iolibs.group_subprocs import IdentifyConfigTag +from madgraph.core.diagram_generation import DiagramTag + +class IdentifyJetTag(IdentifyConfigTag): + """ Like IndentifyConfigTag, but ignores spin and color """ + + @staticmethod + def link_from_leg(leg, model): + (leg_num1, _, mass, width, _), leg_num2 = super( + IdentifyJetTag, IdentifyJetTag + ).link_from_leg(leg, model)[0] + return [((leg_num1, mass, width), leg_num2)] + + @staticmethod + def vertex_id_from_vertex(vertex, last_vertex, model, ninitial): + vertex = super(IdentifyJetTag, IdentifyJetTag).vertex_id_from_vertex( + vertex, last_vertex, model, ninitial + ) + if len(vertex) == 1: + return ((0,),) + (_, mass, width), _ = vertex + return ((mass, width), 0) + + +class IdentifySGJetTag(IdentifySGConfigTag): + """ Like IndentifySGConfigTag, but ignores spin, color and charge """ + + @staticmethod + def link_from_leg(leg, model): + (state, _, _, _, mass, width), leg_num = super( + IdentifySGJetTag, IdentifySGJetTag + ).link_from_leg(leg, model)[0] + return [((state, mass, width), leg_num)] + + @staticmethod + def vertex_id_from_vertex(vertex, last_vertex, model, ninitial): + vertex = super(IdentifySGJetTag, IdentifySGJetTag).vertex_id_from_vertex( + vertex, last_vertex, model, ninitial + ) + if len(vertex) == 1: + return ((0,),) + (_, mass, width, qcd, onshell), = vertex + return ((mass, width, qcd, onshell),) + class OneProcessExporterMG7(export_cpp.OneProcessExporterCPP): @@ -13,15 +57,25 @@ def __init__(self, matrix_element, cpp_helas_call_writer): self.name = f"P{matrix_element.get('processes')[0].shell_string()}" self.model = self.matrix_element.get("processes")[0].get("model") self.amplitude = self.matrix_element.get("base_amplitude") - self.sym_indices, self.sym_perms, _ = find_symmetry( - self.matrix_element, lambda diag: IdentifySGConfigTag(diag, self.model) - ) + merge_jets = False + if merge_jets: + self.sym_indices, self.sym_perms, _ = find_symmetry( + self.matrix_element, + lambda diag: IdentifySGJetTag(diag, self.model), + skip_identical_check=True + ) + else: + self.sym_indices, self.sym_perms, _ = find_symmetry( + self.matrix_element, lambda diag: IdentifySGConfigTag(diag, self.model) + ) + self.diagrams = self.amplitude.get("diagrams") self.helas_diagrams = self.matrix_element.get("diagrams") self.all_flavors, self.all_flavors_pdgs = self.matrix_element.get_external_flavors_with_iden(return_pdgs=True) self.process = self.amplitude.get("process") self.legs = self.process.get("legs_with_decays") self.color_basis = self.matrix_element.get("color_basis") + self.set_subprocess_class() self.set_topology() self.set_flavor_indices() self.set_active_flavors() @@ -30,6 +84,25 @@ def __init__(self, matrix_element, cpp_helas_call_writer): def generate_process_files(self): super().generate_process_files() + def set_subprocess_class(self): + is_parts = [ + self.model.get_particle(l.get("id")) + for l in self.process.get("legs") + if not l.get("state") + ] + fs_parts = [ + self.model.get_particle(l.get("id")) + for l in self.process.get("legs") + if l.get("state") + ] + self.subprocess_class = ( + tuple( + (p.get("mass"), l.get("onshell")) + for (p, l) in zip(is_parts + fs_parts, self.process.get("legs")) + ), + self.process.get("id"), + ) + def set_topology(self): self.edge_names = {} self.incoming = [None] * 2 @@ -76,15 +149,12 @@ def set_channels_colors_map(self): for diag_tuple in self.color_basis[col_basis_elem]: diag_jamps[diag_tuple[0]].append(ijamp) - sym_indices, sym_perms, _ = find_symmetry( - self.matrix_element, lambda diag: IdentifySGConfigTag(diag, self.model) - ) - self.channels = [] - self.channel_indices = [] - for diagram_index, (sym_index, sym_perm) in enumerate(zip(sym_indices, sym_perms)): + channel_indices = [] + self.diagram_tags = [] + for diagram_index, (sym_index, sym_perm) in enumerate(zip(self.sym_indices, self.sym_perms)): if sym_index == 0: - self.channel_indices.append(-1) + channel_indices.append(-1) continue active_colors = diag_jamps[diagram_index] if self.color_basis else [0] @@ -93,8 +163,13 @@ def set_channels_colors_map(self): raise RuntimeError( f"no valid flavor configurations found for diagram {diagram_index+1}" ) + diagram = self.diagrams[diagram_index] if sym_index < 0: - self.channels[self.channel_indices[-sym_index - 1]]["diagrams"].append( + chan_index = channel_indices[-sym_index - 1] + self.diagram_tags[chan_index].append( + IdentifyJetTag(diagram, self.model), + ) + self.channels[chan_index]["diagrams"].append( { "diagram": diagram_index, "permutation": sym_perm, @@ -102,10 +177,9 @@ def set_channels_colors_map(self): "active_colors": active_colors, } ) - self.channel_indices.append(-1) + channel_indices.append(-1) continue - diagram = self.diagrams[diagram_index] vertices = [] propagators = [] on_shell_propagators = [] @@ -129,7 +203,9 @@ def set_channels_colors_map(self): on_shell_propagators.append(prop_index) vertices.append(vertex_props) - self.channel_indices.append(len(self.channels)) + chan_index = len(self.channels) + self.diagram_tags.append([IdentifyJetTag(diagram, self.model)]) + channel_indices.append(chan_index) self.channels.append( { "propagators": propagators, @@ -211,15 +287,20 @@ def get_subprocess_info(self, proc_dir, lib_me_path): } for index, options in self.all_flavors_same_initial ] - return { - "incoming": self.incoming, - "outgoing": self.outgoing, - "channels": self.channels, - "me_path": lib_me_path, - "path": proc_dir, - "flavors": flavors, - "color_flows": color_flows, - "pdg_color_types": pdg_color_types, - "diagram_count": len(self.diagrams), - "helicities": list(self.matrix_element.get_helicity_matrix()), - } + + return ( + { + "incoming": self.incoming, + "outgoing": self.outgoing, + "channels": self.channels, + "me_path": lib_me_path, + "path": proc_dir, + "flavors": flavors, + "color_flows": color_flows, + "pdg_color_types": pdg_color_types, + "diagram_count": len(self.diagrams), + "helicities": list(self.matrix_element.get_helicity_matrix()), + }, + self.diagram_tags, + self.subprocess_class, + ) diff --git a/madgraph/iolibs/template_files/mg7/madevent.py b/madgraph/iolibs/template_files/mg7/madevent.py index 48078090d..0903b9279 100644 --- a/madgraph/iolibs/template_files/mg7/madevent.py +++ b/madgraph/iolibs/template_files/mg7/madevent.py @@ -10,7 +10,7 @@ import subprocess import re import logging -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Literal, NamedTuple import resource @@ -101,8 +101,8 @@ def format_time(t: int, centi: bool = False): class Channel: phasespace_mapping: ms.PhaseSpaceMapping adaptive_mapping: ms.Flow | ms.VegasMapping - discrete_before: ms.DiscreteSampler | ms.DiscreteFlow | None - discrete_after: ms.DiscreteSampler | ms.DiscreteFlow | None + discrete_sym: ms.DiscreteSampler | ms.DiscreteFlow | None + discrete_flavor: ms.DiscreteSampler | ms.DiscreteFlow | None channel_weight_indices: list[int] | None name: str active_flavors: list[int] @@ -114,14 +114,17 @@ class PhaseSpace: mode: Literal["multichannel", "flat", "both"] channels: list[Channel] symfact: list[int | None] - chan_weight_remap: list[int] + first_chan_weight_remap: list[list[int]] = field(default_factory=list) + first_remapped_chan_count: int = 0 + second_chan_weight_remap: list[int] = field(default_factory=list) + second_remapped_chan_count: int = 0 prop_chan_weights: ms.PropagatorChannelWeights | None = None subchan_weights: ms.SubchannelWeights | None = None cwnet: ms.ChannelWeightNetwork | None = None class MultiChannelData(NamedTuple): - amp2_remap: list[int] + amp2_remaps: list[list[int]] symfact: list[int | None] topologies: list[list[ms.Topology]] permutations: list[list[list[int]]] @@ -166,6 +169,12 @@ def load_cards(self) -> None: self.param_card = ParamCard(self.param_card_path) with open(os.path.join("SubProcesses", "subprocesses.json")) as f: self.subprocess_data = json.load(f) + if self.run_card["phasespace"]["merge_subprocesses"]: + with open(os.path.join("SubProcesses", "merged_subprocesses.json")) as f: + self.merged_subprocess_data = json.load(f) + else: + self.merged_subprocess_data = None + def init_backend(self) -> None: ms.set_simd_vector_size(self.run_card["run"]["simd_vector_size"]) @@ -381,8 +390,14 @@ def init_generator_config(self) -> None: def init_subprocesses(self) -> None: self.subprocesses = [] - for subproc_id, meta in enumerate(self.subprocess_data): - self.subprocesses.append(MadgraphSubprocess(self, meta, subproc_id)) + if self.merged_subprocess_data is None: + for subproc_id, meta in enumerate(self.subprocess_data): + self.subprocesses.append(MadgraphSubprocess(self, meta, subproc_id)) + else: + for subproc_id, meta in enumerate(self.merged_subprocess_data): + self.subprocesses.append( + MadgraphSubprocess(self, meta, subproc_id, self.subprocess_data) + ) def build_event_generator(self, phasespaces: list[PhaseSpace]) -> ms.EventGenerator: channel_generators = [] @@ -448,14 +463,12 @@ def survey(self) -> None: for subproc in self.subprocesses ] evgen_multi = self.survey_phasespaces(phasespaces_multi) - phasespaces_flat = [ subproc.build_flat_phasespace() - if len(subproc.meta["channels"]) > kept_count + 1 else + if len(ps_multi.channels) > kept_count + 1 else None - for subproc in self.subprocesses + for subproc, ps_multi in zip(self.subprocesses, phasespaces_multi) ] - #evgen_flat = self.survey_phasespaces(phasespaces_flat, "flat") channel_status = evgen_multi.channel_status() cross_sections = [] @@ -731,37 +744,30 @@ def get_result(self) -> dict: return result def build_lhe_completer(self): - subproc_args = [] - for subproc, meta in zip(self.subprocesses, self.subprocess_data): - ( - _, - _, - topologies, - permutations, - _, - _, - diagram_indices, - diagram_color_indices, - _, - ) = subproc.build_multi_channel_data() - subproc_args.append( - ms.SubprocArgs( - topologies = [topo[0] for topo in topologies], - permutations = permutations, - diagram_indices = diagram_indices, - diagram_color_indices = diagram_color_indices, - color_flows = meta["color_flows"], - pdg_color_types = { - int(key): value - for key, value in meta["pdg_color_types"].items() - }, - helicities = meta["helicities"], - pdg_ids = [flavor["options"] for flavor in meta["flavors"]], - matrix_flavor_indices = [ - flavor["index"] for flavor in meta["flavors"] - ], - ) + all_mcdata = ( + [subproc.build_multi_channel_data() for subproc in self.subprocesses] + if self.merged_subprocess_data is None else + [build_multi_channel_data(meta, self) for meta in self.subprocess_data] + ) + subproc_args = [ + ms.SubprocArgs( + topologies = [topo[0] for topo in mcdata.topologies], + permutations = mcdata.permutations, + diagram_indices = mcdata.diagram_indices, + diagram_color_indices = mcdata.diagram_color_indices, + color_flows = meta["color_flows"], + pdg_color_types = { + int(key): value + for key, value in meta["pdg_color_types"].items() + }, + helicities = meta["helicities"], + pdg_ids = [flavor["options"] for flavor in meta["flavors"]], + matrix_flavor_indices = [ + flavor["index"] for flavor in meta["flavors"] + ], ) + for mcdata, meta in zip(all_mcdata, self.subprocess_data) + ] return ms.LHECompleter( subproc_args=subproc_args, bw_cutoff=self.run_card["phasespace"]["bw_cutoff"] @@ -871,36 +877,176 @@ def clean_pids(pids: list[int]) -> list[int]: return pids_out +def build_topologies( + incoming_masses: list[float], + outgoing_masses: list[float], + channel: dict, + process: MadgraphProcess +) -> list[ms.Topology]: + propagators = [] + for i, pid in enumerate(clean_pids(channel["propagators"])): + mass = process.get_mass(pid) + width = process.get_width(pid) + if i in channel["on_shell_propagators"]: + bw_cutoff = process.run_card["phasespace"]["bw_cutoff"] + e_min = mass - bw_cutoff * width + e_max = mass + bw_cutoff * width + else: + e_min = 0 + e_max = 0 + propagators.append(ms.Propagator( + mass=mass, + width=width, + integration_order=0, + e_min=e_min, + e_max=e_max, + )) + vertices = channel["vertices"] + diag = ms.Diagram( + incoming_masses, outgoing_masses, propagators, vertices + ) + return ms.Topology.topologies(diag) + + +def build_multi_channel_data( + meta: dict, process: MadgraphProcess, unmerged_meta: dict | None = None +) -> MultiChannelData: + incoming_masses = [ + process.get_mass(pid) for pid in clean_pids(meta["incoming"]) + ] + outgoing_masses = [ + process.get_mass(pid) for pid in clean_pids(meta["outgoing"]) + ] + + if unmerged_meta is None: + diagram_count = meta["diagram_count"] + amp2_remaps = [[-1] * diagram_count] + else: + amp2_remaps = [ + [-1] * unmerged_meta[subproc]["diagram_count"] + for subproc in meta["subprocesses"] + ] + symfact = [] + topologies = [] + permutations = [] + channel_indices = [] + channel_weight_indices = [] + diagram_indices = [] + diagram_color_indices = [] + active_flavors = [] + channel_index = 0 + + for channel in meta["channels"]: + if unmerged_meta is None: + topo_channel = channel + else: + topo_subproc = channel["subprocess"] + topo_channel_index = channel["channel"] + topo_channel = unmerged_meta[topo_subproc]["channels"][topo_channel_index] + chan_topologies = build_topologies( + incoming_masses, outgoing_masses, topo_channel, process + ) + topo_count = len(chan_topologies) + diagrams = channel["diagrams"] + chan_permutations = [d["permutation"] for d in diagrams] + if unmerged_meta is None: + amp2_remaps[0][diagrams[0]["diagram"]] = channel_index + else: + for amp2_remap, diag in zip(amp2_remaps, diagrams[0]["diagram"]): + if diag != -1: + amp2_remap[diag] = channel_index + + channel_index_first = channel_index + symfact_index_first = len(symfact) + channel_index += 1 + symfact.extend([None] * topo_count) + for d in diagrams[1:]: + if unmerged_meta is None: + amp2_remaps[0][d["diagram"]] = channel_index + else: + for amp2_remap, diag in zip(amp2_remaps, d["diagram"]): + if diag != -1: + amp2_remap[diag] = channel_index + channel_index += 1 + symfact.extend(range(symfact_index_first, symfact_index_first + topo_count)) + + topologies.append(chan_topologies) + permutations.append(chan_permutations) + channel_indices.append(list(range(channel_index_first, channel_index))) + channel_weight_indices.append([ + [ + symfact_index_first + topo_index + i * topo_count + for i in range(len(chan_permutations)) + ] + for topo_index in range(topo_count) + ]) + diagram_indices.append([d["diagram"] for d in diagrams]) + if unmerged_meta is None: + diagram_color_indices.append([d["active_colors"] for d in diagrams]) + active_flavors.append([d["active_flavors"] for d in diagrams]) + + return MultiChannelData( + amp2_remaps, + symfact, + topologies, + permutations, + channel_indices, + channel_weight_indices, + diagram_indices, + diagram_color_indices, + active_flavors, + ) + + class MadgraphSubprocess: - def __init__(self, process: MadgraphProcess, meta: dict, subproc_id: int): + def __init__( + self, + process: MadgraphProcess, + meta: dict, + subproc_id: int, + unmerged_meta: dict | None = None + ): self.process = process self.meta = meta + self.unmerged_meta = unmerged_meta self.subproc_id = subproc_id self.multi_channel_data = None - api_path_format = self.meta["me_path"] - subproc_path = self.meta["path"] - devices = self.process.run_card["run"]["devices"] - api_paths = [] - if not isinstance(devices, list): - devices = [devices] - for device in devices: - subproc_dir = os.path.dirname(subproc_path) - # 'cppauto' resolve quick fix - resolved = device - if device == "cppauto": - out = subprocess.run( - ["make", "-n", "BACKEND=cppauto", "detect-backend"], - cwd=subproc_path, capture_output=True, text=True, - ).stdout - match = re.search(r"BACKEND=(\S+) \(was cppauto\)", out) - if match: - resolved = match.group(1) - api_path = api_path_format.format(device=resolved) - if not os.path.isfile(api_path): - logger.info(f"Compiling subprocess {subproc_dir}, for device '{device}'") - misc.compile(arg = [f"BACKEND={device}", "USEBUILDDIR=1"], cwd = subproc_path) - api_paths.append(api_path) + if unmerged_meta is None: + api_path_formats = [self.meta["me_path"]] + subproc_paths = [self.meta["path"]] + else: + api_path_formats = [] + subproc_paths = [] + for subproc in self.meta["subprocesses"]: + submeta = self.unmerged_meta[subproc] + api_path_formats.append(submeta["me_path"]) + subproc_paths.append(submeta["path"]) + + all_api_paths = [] + for api_path_format, subproc_path in zip(api_path_formats, subproc_paths): + devices = self.process.run_card["run"]["devices"] + api_paths = [] + if not isinstance(devices, list): + devices = [devices] + for device in devices: + subproc_dir = os.path.dirname(subproc_path) + # 'cppauto' resolve quick fix + resolved = device + if device == "cppauto": + out = subprocess.run( + ["make", "-n", "BACKEND=cppauto", "detect-backend"], + cwd=subproc_path, capture_output=True, text=True, + ).stdout + match = re.search(r"BACKEND=(\S+) \(was cppauto\)", out) + if match: + resolved = match.group(1) + api_path = api_path_format.format(device=resolved) + if not os.path.isfile(api_path): + logger.info(f"Compiling subprocess {subproc_dir}, for device '{device}'") + misc.compile(arg = [f"BACKEND={device}", "USEBUILDDIR=1"], cwd = subproc_path) + api_paths.append(api_path) + all_api_paths.append(api_paths) self.incoming_masses = [ self.process.get_mass(pid) for pid in clean_pids(self.meta["incoming"]) @@ -942,103 +1088,33 @@ def __init__(self, process: MadgraphProcess, meta: dict, subproc_id: int): ) if self.process.run_card["run"]["dummy_matrix_element"]: - self.matrix_element = None + self.matrix_elements = [None] * len(all_api_paths) else: - for context, api_path in zip(self.process.contexts, api_paths): - self.matrix_element = context.load_matrix_element( - api_path, self.process.param_card_path - ) + self.matrix_elements = [] + for api_paths in all_api_paths: + for context, api_path in zip(self.process.contexts, api_paths): + mat = context.load_matrix_element( + api_path, self.process.param_card_path + ) + self.matrix_elements.append(mat) def build_multi_channel_data(self) -> MultiChannelData: if self.multi_channel_data is not None: return self.multi_channel_data - - diagram_count = self.meta["diagram_count"] - bw_cutoff = self.process.run_card["phasespace"]["bw_cutoff"] - - amp2_remap = [-1] * diagram_count - symfact = [] - topologies = [] - permutations = [] - channel_indices = [] - channel_weight_indices = [] - diagram_indices = [] - diagram_color_indices = [] - active_flavors = [] - channel_index = 0 - - for channel_id, channel in enumerate(self.meta["channels"]): - propagators = [] - for i, pid in enumerate(clean_pids(channel["propagators"])): - mass = self.process.get_mass(pid) - width = self.process.get_width(pid) - if i in channel["on_shell_propagators"]: - e_min = mass - bw_cutoff * width - e_max = mass + bw_cutoff * width - else: - e_min = 0 - e_max = 0 - propagators.append(ms.Propagator( - mass=mass, - width=width, - integration_order=0, - e_min=e_min, - e_max=e_max, - )) - vertices = channel["vertices"] - diagrams = channel["diagrams"] - chan_permutations = [d["permutation"] for d in diagrams] - diag = ms.Diagram( - self.incoming_masses, self.outgoing_masses, propagators, vertices - ) - chan_topologies = ms.Topology.topologies(diag) - topo_count = len(chan_topologies) - - amp2_remap[diagrams[0]["diagram"]] = channel_index - channel_index_first = channel_index - symfact_index_first = len(symfact) - channel_index += 1 - symfact.extend([None] * topo_count) - for d in diagrams[1:]: - amp2_remap[d["diagram"]] = channel_index - channel_index += 1 - symfact.extend(range(symfact_index_first, symfact_index_first + topo_count)) - - topologies.append(chan_topologies) - permutations.append(chan_permutations) - channel_indices.append(list(range(channel_index_first, channel_index))) - channel_weight_indices.append([ - [ - symfact_index_first + topo_index + i * topo_count - for i in range(len(chan_permutations)) - ] - for topo_index in range(topo_count) - ]) - diagram_indices.append([d["diagram"] for d in diagrams]) - diagram_color_indices.append([d["active_colors"] for d in diagrams]) - active_flavors.append([d["active_flavors"] for d in diagrams]) - self.multi_channel_data = MultiChannelData( - amp2_remap, - symfact, - topologies, - permutations, - channel_indices, - channel_weight_indices, - diagram_indices, - diagram_color_indices, - active_flavors, + self.multi_channel_data = build_multi_channel_data( + self.meta, self.process, self.unmerged_meta ) return self.multi_channel_data def build_multichannel_phasespace(self) -> PhaseSpace: ( - amp2_remap, + amp2_remaps, symfact, topologies, permutations, channel_indices, channel_weight_indices, - diagram_indices, + _, _, all_active_flavors, ) = self.build_multi_channel_data() @@ -1064,44 +1140,44 @@ def build_multichannel_phasespace(self) -> PhaseSpace: prefix = f"subproc{self.subproc_id}.channel{channel_id}" if topo_count > 1: prefix += f".subchan{topo_index}" - discrete_before, discrete_after = self.build_discrete( + discrete_sym, discrete_flavor = self.build_discrete( len(chan_permutations), len(self.meta["flavors"]), prefix ) channels.append(Channel( phasespace_mapping = mapping, adaptive_mapping = self.build_vegas(mapping, prefix), - discrete_before = discrete_before, - discrete_after = discrete_after, + discrete_sym = discrete_sym, + discrete_flavor = discrete_flavor, channel_weight_indices = indices, name = f"{channel_id}", active_flavors = active_flavors, )) - chan_weight_remap = list(range(len(symfact))) #TODO: only construct if necessary + remapped_chan_count = sum(len(indices) for indices in channel_indices) if self.process.run_card["phasespace"]["sde_strategy"] == "denominators": prop_chan_weights = ms.PropagatorChannelWeights( [topo[0] for topo in topologies], permutations, channel_indices ) - indices_for_subchan = channel_indices + chan_weight_remap = [] else: prop_chan_weights = None - indices_for_subchan = diagram_indices + chan_weight_remap = [ + [len(symfact) if remap == -1 else remap for remap in amp2_remap] + for amp2_remap in amp2_remaps + ] if any(len(topos) > 1 for topos in topologies): subchan_weights = ms.SubchannelWeights( - topologies, permutations, indices_for_subchan + topologies, permutations, channel_indices ) else: subchan_weights = None - if prop_chan_weights is None: - chan_weight_remap = [ - len(symfact) if remap == -1 else remap for remap in amp2_remap - ] return PhaseSpace( mode="multichannel", channels=channels, - chan_weight_remap=chan_weight_remap, + first_chan_weight_remap=chan_weight_remap, + first_remapped_chan_count=remapped_chan_count, symfact=symfact, prop_chan_weights=prop_chan_weights, subchan_weights=subchan_weights, @@ -1116,22 +1192,30 @@ def build_flat_phasespace(self) -> PhaseSpace: leptonic=self.process.leptonic, ) prefix = f"subproc{self.subproc_id}.flat" - discrete_before, discrete_after = self.build_discrete( + discrete_sym, discrete_flavor = self.build_discrete( 1, len(self.meta["flavors"]), prefix ) channel = Channel( phasespace_mapping = mapping, adaptive_mapping = self.build_vegas(mapping, prefix), - discrete_before = discrete_before, - discrete_after = discrete_after, + discrete_sym = discrete_sym, + discrete_flavor = discrete_flavor, channel_weight_indices = [0], name = "F", active_flavors = [], ) + if self.unmerged_meta is None: + remap = [list(range(self.meta["diagram_count"]))] + else: + remap = [ + list(range(self.unmerged_meta[subproc]["diagram_count"])) + for subproc in self.meta["subprocesses"] + ] return PhaseSpace( mode="flat", channels=[channel], - chan_weight_remap=[0] * self.meta["diagram_count"], + first_chan_weight_remap=remap, + first_remapped_chan_count=1, symfact=[None], ) @@ -1176,8 +1260,8 @@ def simplify_phasespace( channels.append(Channel( phasespace_mapping = channel.phasespace_mapping, adaptive_mapping = channel.adaptive_mapping, - discrete_before = channel.discrete_before, - discrete_after = channel.discrete_after, + discrete_sym = channel.discrete_sym, + discrete_flavor = channel.discrete_flavor, channel_weight_indices = list(range( channel_index, channel_index + perm_count )), @@ -1190,8 +1274,8 @@ def simplify_phasespace( channels.append(Channel( phasespace_mapping = flat_channel.phasespace_mapping, adaptive_mapping = flat_channel.adaptive_mapping, - discrete_before = flat_channel.discrete_before, - discrete_after = flat_channel.discrete_after, + discrete_sym = flat_channel.discrete_sym, + discrete_flavor = flat_channel.discrete_flavor, channel_weight_indices = [len(symfact)], name = flat_channel.name, active_flavors = flat_channel.active_flavors, @@ -1199,15 +1283,34 @@ def simplify_phasespace( flat_index = len(symfact) symfact.append(None) channel_map[len(multi_phasespace.symfact)] = len(symfact) - chan_weight_remap = [ - channel_map.get(remap, flat_index) - for remap in multi_phasespace.chan_weight_remap - ] + if multi_phasespace.subchan_weights is None and len(multi_phasespace.first_chan_weight_remap) > 0: + first_chan_weight_remap = [ + [ + channel_map.get(remap, flat_index) + for remap in cw_remap + ] + for cw_remap in multi_phasespace.first_chan_weight_remap + ] + first_remapped_chan_count = len(symfact) + second_chan_weight_remap = [] + second_remapped_chan_count = 0 + else: + first_chan_weight_remap = multi_phasespace.first_chan_weight_remap + first_remapped_chan_count = multi_phasespace.first_remapped_chan_count + chan_count = multi_phasespace.first_remapped_chan_count if multi_phasespace.subchan_weights is None else multi_phasespace.subchan_weights.channel_count() + second_chan_weight_remap = [ + channel_map.get(i, flat_index) + for i in range(chan_count) + ] + second_remapped_chan_count = len(symfact) return PhaseSpace( mode="both", channels=channels, - chan_weight_remap=chan_weight_remap, + first_chan_weight_remap=first_chan_weight_remap, + first_remapped_chan_count=first_remapped_chan_count, + second_chan_weight_remap=second_chan_weight_remap, + second_remapped_chan_count=second_remapped_chan_count, symfact=symfact, prop_chan_weights=multi_phasespace.prop_chan_weights, subchan_weights=multi_phasespace.subchan_weights, @@ -1220,21 +1323,7 @@ def build_madnis(self, phasespace: PhaseSpace) -> PhaseSpace: prefix = f"subproc{self.subproc_id}.channel{channel_id}" cond_dim = 0 - discrete_before = channel.discrete_before - if discrete_before is not None: - perm_count = channel.phasespace_mapping.channel_count() - discrete_before = ms.DiscreteFlow( - option_counts=[perm_count], - prefix=f"{prefix}.discrete_flow_before", - dims_with_prior=[], - condition_dim=0, - subnet_hidden_dim=madnis_args["discrete_hidden_dim"], - subnet_layers=madnis_args["discrete_layers"], - subnet_activation=self.activation(madnis_args["discrete_activation"]), - ) - discrete_before.initialize_globals(self.process.contexts[0]) - cond_dim += perm_count - + # The adaptive map runs before discrete_sym, so it is unconditioned. flow_dim = channel.phasespace_mapping.random_dim() flow = ms.Flow( input_dim=flow_dim, @@ -1254,24 +1343,40 @@ def build_madnis(self, phasespace: PhaseSpace) -> PhaseSpace: ) cond_dim += flow_dim - discrete_after = channel.discrete_after - if discrete_after is not None: - discrete_after = ms.DiscreteFlow( + # discrete_sym runs after the adaptive map, so it can condition on its latent. + discrete_sym = channel.discrete_sym + if discrete_sym is not None: + perm_count = channel.phasespace_mapping.channel_count() + discrete_sym = ms.DiscreteFlow( + option_counts=[perm_count], + prefix=f"{prefix}.discrete_flow_sym", + dims_with_prior=[], + condition_dim=cond_dim, + subnet_hidden_dim=madnis_args["discrete_hidden_dim"], + subnet_layers=madnis_args["discrete_layers"], + subnet_activation=self.activation(madnis_args["discrete_activation"]), + ) + discrete_sym.initialize_globals(self.process.contexts[0]) + cond_dim += perm_count + + discrete_flavor = channel.discrete_flavor + if discrete_flavor is not None: + discrete_flavor = ms.DiscreteFlow( option_counts=[len(self.meta["flavors"])], - prefix=f"{prefix}.discrete_flow_after", + prefix=f"{prefix}.discrete_flow_flavor", dims_with_prior=[0], condition_dim=cond_dim, subnet_hidden_dim=madnis_args["discrete_hidden_dim"], subnet_layers=madnis_args["discrete_layers"], subnet_activation=self.activation(madnis_args["discrete_activation"]), ) - discrete_after.initialize_globals(self.process.contexts[0]) + discrete_flavor.initialize_globals(self.process.contexts[0]) channels.append(Channel( phasespace_mapping = channel.phasespace_mapping, adaptive_mapping = flow, - discrete_before = discrete_before, - discrete_after = discrete_after, + discrete_sym = discrete_sym, + discrete_flavor = discrete_flavor, channel_weight_indices = channel.channel_weight_indices, name = channel.name, active_flavors = channel.active_flavors, @@ -1280,7 +1385,10 @@ def build_madnis(self, phasespace: PhaseSpace) -> PhaseSpace: return PhaseSpace( mode="both", channels=channels, - chan_weight_remap=phasespace.chan_weight_remap, + first_chan_weight_remap=phasespace.first_chan_weight_remap, + first_remapped_chan_count=phasespace.first_remapped_chan_count, + second_chan_weight_remap=phasespace.second_chan_weight_remap, + second_remapped_chan_count=phasespace.second_remapped_chan_count, symfact=phasespace.symfact, cwnet=self.build_cwnet(len(phasespace.symfact)), prop_chan_weights=phasespace.prop_chan_weights, @@ -1303,28 +1411,30 @@ def build_vegas(self, mapping: ms.PhaseSpaceMapping, prefix: str) -> ms.VegasMap def build_discrete( self, permutation_count: int, flavor_count: int, prefix: str ) -> tuple[ms.DiscreteSampler | None, ms.DiscreteSampler | None]: - discrete_before = None - #if permutation_count > 1: - # discrete_before = ms.DiscreteSampler( - # [permutation_count], f"{prefix}.discrete_before" - # ) - # for context in self.process.contexts: - # discrete_before.initialize_globals(context) - #else: - # discrete_before = None + is_adaptive = self.process.run_card["phasespace"]["adaptive_symmetry_sampling"] + if is_adaptive and permutation_count > 1: + discrete_sym = ms.DiscreteSampler( + [permutation_count], f"{prefix}.discrete_sym" + ) + for context in self.process.contexts: + discrete_sym.initialize_globals(context) + else: + discrete_sym = None if flavor_count > 1: - discrete_after = ms.DiscreteSampler( - [flavor_count], f"{prefix}.discrete_after", [0] + discrete_flavor = ms.DiscreteSampler( + [flavor_count], f"{prefix}.discrete_flavor", [0] ) for context in self.process.contexts: - discrete_after.initialize_globals(context) + discrete_flavor.initialize_globals(context) else: - discrete_after = None + discrete_flavor = None - return discrete_before, discrete_after + return discrete_sym, discrete_flavor def build_cwnet(self, channel_count: int) -> ms.ChannelWeightNetwork: + #if channel_count == 1: + # return None madnis_args = self.process.run_card["madnis"] cwnet = ms.ChannelWeightNetwork( channel_count=channel_count, @@ -1372,56 +1482,77 @@ def build_integrands( flavor_remap = [] flavor_factors = [] flavor_mirror = [] + flavor_diff_xs_indices = [] + flavor_subproc_indices = [] + flavor_per_subproc_remap = [] + for flav in self.meta["flavors"]: + if self.unmerged_meta is not None: + diff_xs_index = flav["subprocess"] + subproc_index = self.meta["subprocesses"][diff_xs_index] + ps_flavor = flav["flavor"] + flavor_diff_xs_indices.append(diff_xs_index) + flavor_subproc_indices.append(subproc_index) + flavor_per_subproc_remap.append(ps_flavor) + flav = self.unmerged_meta[subproc_index]["flavors"][ps_flavor] flavors.append(flav["options"][0]) 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, + + cross_sections = [] + for matrix_element in self.matrix_elements: + if matrix_element: + mat = ms.MatrixElement( + matrix_element, + ms.Integrand.matrix_element_inputs, + ms.Integrand.matrix_element_outputs, + True, + ) + else: + #TODO: not working in merged mode + mat = ms.MatrixElement( + 0xBADCAFE, + self.particle_count, + ms.Integrand.matrix_element_inputs, + ms.Integrand.matrix_element_outputs, + self.meta["diagram_count"], + True, + ) + pdf_grid = None if self.process.leptonic else self.process.pdf_grid + pdf_arg = None if self.process.leptonic else ms.CachedPdf() + cross_sections.append( + ms.DifferentialCrossSection( + matrix_element=mat, + cm_energy=self.process.e_cm, + running_coupling=None, + energy_scale=ms.CachedScale(), + pid_options=[], + pdf1=pdf_arg, + pdf2=pdf_arg, + input_momentum_fraction=True, + ) ) - 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, - ) partial_weights = self.process.run_card["generation"]["systematics"] integrands = [] for channel in phasespace.channels: integrands.append(ms.Integrand( channel.phasespace_mapping, - cross_section, + cross_sections, channel.adaptive_mapping, - channel.discrete_before, - channel.discrete_after, + channel.discrete_sym, + channel.discrete_flavor, + flavors, pdf_grid, self.process.running_coupling, self.scale, phasespace.prop_chan_weights, phasespace.subchan_weights, phasespace.cwnet, - phasespace.chan_weight_remap, - len(phasespace.symfact), + phasespace.first_chan_weight_remap, + phasespace.first_remapped_chan_count, + phasespace.second_chan_weight_remap, + phasespace.second_remapped_chan_count, madnis_training, drop_cuts_and_rescale, partial_weights, @@ -1430,8 +1561,11 @@ def build_integrands( flavor_remap, flavor_factors, flavor_mirror, + flavor_diff_xs_indices, + flavor_subproc_indices, + flavor_per_subproc_remap, )) - #print(integrands[0].function()) + #print(integrands[1].function()) #for i in integrands: print(i.function()) return integrands diff --git a/madgraph/iolibs/template_files/mg7/run_card.toml b/madgraph/iolibs/template_files/mg7/run_card.toml index c48d95a2c..7a0a98947 100644 --- a/madgraph/iolibs/template_files/mg7/run_card.toml +++ b/madgraph/iolibs/template_files/mg7/run_card.toml @@ -57,6 +57,7 @@ max_batch_size = %(vegas.max_batch_size)s [phasespace] mode = %(phasespace.mode)s #options: multichannel, flat, both +merge_subprocesses = %(phasespace.merge_subprocesses)s sde_strategy = %(phasespace.sde_strategy)s #options: diagrams, denominators decays = %(phasespace.decays)s # options: all, massive, none t_channel = %(phasespace.t_channel)s # options: propagator, rambo, chili @@ -64,6 +65,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 +adaptive_symmetry_sampling = %(phasespace.adaptive_symmetry_sampling)s [multiparticles] $multiparticles diff --git a/madgraph/various/banner.py b/madgraph/various/banner.py index 2120ed298..35ab72edf 100755 --- a/madgraph/various/banner.py +++ b/madgraph/various/banner.py @@ -6494,6 +6494,7 @@ def default_setup(self): # -------------------------- [phasespace] ---------------------- self.add_toml_param('phasespace', 'mode', "multichannel", allowed=['multichannel', 'flat', 'both']) + self.add_toml_param('phasespace', 'merge_subprocesses', False) self.add_toml_param('phasespace', 'sde_strategy', "diagrams", allowed=['diagrams', 'denominators']) self.add_toml_param('phasespace', 'decays', "all", @@ -6505,6 +6506,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', 'adaptive_symmetry_sampling', True) # ----------------------------- [madnis] ----------------------- self.add_toml_param('madnis', 'enable', False) diff --git a/madgraph/various/diagram_symmetry.py b/madgraph/various/diagram_symmetry.py index 1cb7baea7..30680be76 100755 --- a/madgraph/various/diagram_symmetry.py +++ b/madgraph/various/diagram_symmetry.py @@ -67,7 +67,8 @@ # find_symmetry #=============================================================================== -def find_symmetry(matrix_element, tag_type=diagram_generation.DiagramTag): +def find_symmetry(matrix_element, tag_type=diagram_generation.DiagramTag, + skip_identical_check=False): """Find symmetries between amplitudes by comparing diagram tags for all the diagrams in the process. Identical diagram tags correspond to different external particle permutations of the same @@ -123,7 +124,7 @@ def find_symmetry(matrix_element, tag_type=diagram_generation.DiagramTag): symmetry.append(1) # Check for matrix elements with no identical particles - if matrix_element.get("identical_particle_factor") == 1: + if not skip_identical_check and matrix_element.get("identical_particle_factor") == 1: return symmetry, \ permutations,\ [list(range(nexternal))] @@ -144,7 +145,6 @@ def find_symmetry(matrix_element, tag_type=diagram_generation.DiagramTag): # Only 3-vertices allowed in configs.inc continue - #tag = diagram_generation.DiagramTag(base_diagram) tag = tag_type(base_diagram) try: ind = diagram_tags.index(tag) diff --git a/madspace/include/madspace/compgraphs/function_builder_mixin.inc b/madspace/include/madspace/compgraphs/function_builder_mixin.inc index 360d8db14..3455951e5 100644 --- a/madspace/include/madspace/compgraphs/function_builder_mixin.inc +++ b/madspace/include/madspace/compgraphs/function_builder_mixin.inc @@ -59,6 +59,14 @@ Value accept_norm(Value accepted_batch, Value full_batch) { return instruction("accept_norm", {accepted_batch, full_batch})[0]; } +ValueVec batch_split_by_index(Value indices, Value count) { + return instruction("batch_split_by_index", {indices, count}); +} + +Value batch_merge_by_index(ValueVec args) { + return instruction("batch_merge_by_index", args)[0]; +} + Value add(Value in1, Value in2) { return instruction("add", {in1, in2})[0]; } diff --git a/madspace/include/madspace/compgraphs/instruction.hpp b/madspace/include/madspace/compgraphs/instruction.hpp index 184c83948..8cdc7c8d7 100644 --- a/madspace/include/madspace/compgraphs/instruction.hpp +++ b/madspace/include/madspace/compgraphs/instruction.hpp @@ -115,6 +115,20 @@ class BatchSplitInstruction : public Instruction { TypeVec signature(const ValueVec& args) const override; }; +class BatchSplitByIndexInstruction : public Instruction { +public: + BatchSplitByIndexInstruction(int opcode, bool differentiable) : + Instruction("batch_split_by_index", opcode, differentiable) {} + TypeVec signature(const ValueVec& args) const override; +}; + +class BatchMergeByIndexInstruction : public Instruction { +public: + BatchMergeByIndexInstruction(int opcode, bool differentiable) : + Instruction("batch_merge_by_index", opcode, differentiable) {} + TypeVec signature(const ValueVec& args) const override; +}; + class CatInstruction : public Instruction { public: CatInstruction(int opcode, bool differentiable) : diff --git a/madspace/include/madspace/compgraphs/opcode_mixin.inc b/madspace/include/madspace/compgraphs/opcode_mixin.inc index ba664afd0..8dbebab93 100644 --- a/madspace/include/madspace/compgraphs/opcode_mixin.inc +++ b/madspace/include/madspace/compgraphs/opcode_mixin.inc @@ -12,146 +12,148 @@ full = 10, squeeze = 11, unsqueeze = 12, 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 +batch_split_by_index = 14, +batch_merge_by_index = 15, +add = 16, +add_int = 17, +sub = 18, +mul = 19, +div = 20, +reduce_sum = 21, +reduce_sum_vector = 22, +batch_reduce_mean = 23, +batch_reduce_mean_keepdim = 24, +reduce_product = 25, +sqrt = 26, +square = 27, +min = 28, +max = 29, +obs_sqrt_s = 30, +obs_e = 31, +obs_px = 32, +obs_py = 33, +obs_pz = 34, +obs_mass = 35, +obs_pt = 36, +obs_p_mag = 37, +obs_phi = 38, +obs_theta = 39, +obs_y = 40, +obs_y_abs = 41, +obs_eta = 42, +obs_eta_abs = 43, +obs_delta_eta = 44, +obs_delta_phi = 45, +obs_delta_r = 46, +boost_beam = 47, +boost_beam_inverse = 48, +com_p_in = 49, +r_to_x1x2 = 50, +x1x2_to_r = 51, +diff_cross_section = 52, +two_body_decay_com = 53, +two_body_decay_com_inverse = 54, +two_body_decay = 55, +two_body_decay_inverse = 56, +two_to_two_particle_scattering_com = 57, +two_to_two_particle_scattering_com_inverse = 58, +two_to_two_particle_scattering = 59, +two_to_two_particle_scattering_inverse = 60, +two_to_three_particle_scattering = 61, +two_to_three_particle_scattering_inverse = 62, +double_t_scattering = 63, +double_t_scattering_inverse = 64, +three_body_decay_com = 65, +three_body_decay_com_inverse = 66, +three_body_decay = 67, +three_body_decay_inverse = 68, +t_inv_min_max = 69, +t_inv_value_and_min_max = 70, +t_inv_min_max_cut = 71, +t_inv_value_and_min_max_cut = 72, +t1_inv_min_max_doublet = 73, +t1_inv_value_and_min_max_doublet = 74, +t2_inv_min_max_doublet = 75, +t2_inv_value_and_min_max_doublet = 76, +s23_min_max = 77, +s23_value_and_min_max = 78, +s23_min_max_cut = 79, +s23_value_and_min_max_cut = 80, +invariants_from_momenta = 81, +sde2_channel_weights = 82, +subchannel_weights = 83, +apply_subchannel_weights = 84, +pt_eta_phi_x = 85, +mirror_momenta = 86, +momenta_to_x1x2 = 87, +uniform_invariant = 88, +uniform_invariant_inverse = 89, +breit_wigner_invariant = 90, +breit_wigner_invariant_inverse = 91, +stable_invariant = 92, +stable_invariant_inverse = 93, +stable_invariant_nu = 94, +stable_invariant_nu_inverse = 95, +fast_rambo_massless = 96, +fast_rambo_massless_inverse = 97, +fast_rambo_massless_com = 98, +fast_rambo_massive = 99, +fast_rambo_massive_inverse = 100, +fast_rambo_massive_com = 101, +cut_unphysical = 102, +cut_one = 103, +cut_all = 104, +cut_any = 105, +scale_transverse_energy = 106, +scale_transverse_mass = 107, +scale_half_transverse_mass = 108, +scale_partonic_energy = 109, +chili_forward = 110, +chili_inverse = 111, +matrix_element = 112, +collect_channel_weights = 113, +interpolate_pdf = 114, +interpolate_alpha_s = 115, +matmul = 116, +relu = 117, +leaky_relu = 118, +elu = 119, +gelu = 120, +sigmoid = 121, +softplus = 122, +rqs_reshape = 123, +rqs_find_bin = 124, +rqs_forward = 125, +rqs_inverse = 126, +softmax = 127, +softmax_prior = 128, +sample_discrete = 129, +sample_discrete_inverse = 130, +sample_discrete_probs = 131, +sample_discrete_probs_inverse = 132, +discrete_histogram = 133, +permute_momenta = 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/driver/channel_generator.hpp b/madspace/include/madspace/driver/channel_generator.hpp index 957037d15..f1eab49a4 100644 --- a/madspace/include/madspace/driver/channel_generator.hpp +++ b/madspace/include/madspace/driver/channel_generator.hpp @@ -104,6 +104,7 @@ class ChannelEventGenerator { int color_index, helicity_index, diagram_index, flavor_index; int ren_scale, alpha_qcd; int x1, fact_scale1, x2, fact_scale2, partial_weight_product; + int subprocess_index; int random, rest; }; diff --git a/madspace/include/madspace/phasespace/integrand.hpp b/madspace/include/madspace/phasespace/integrand.hpp index 9a00bc838..88c843fe5 100644 --- a/madspace/include/madspace/phasespace/integrand.hpp +++ b/madspace/include/madspace/phasespace/integrand.hpp @@ -41,18 +41,21 @@ class Integrand : public FunctionGenerator { Integrand( const PhaseSpaceMapping& mapping, - const DifferentialCrossSection& diff_xs, + const std::vector& diff_xs, const AdaptiveMapping& adaptive_map = std::monostate{}, - const AdaptiveDiscrete& discrete_before = std::monostate{}, - const AdaptiveDiscrete& discrete_after = std::monostate{}, + const AdaptiveDiscrete& discrete_sym = std::monostate{}, + const AdaptiveDiscrete& discrete_flavor = std::monostate{}, + const nested_vector2& pid_options = {}, const std::optional& pdf_grid = std::nullopt, const std::optional& running_coupling = std::nullopt, const std::optional& energy_scale = std::nullopt, const std::optional& prop_chan_weights = std::nullopt, const std::optional& subchan_weights = std::nullopt, const std::optional& chan_weight_net = std::nullopt, - const std::vector& chan_weight_remap = {}, - std::size_t remapped_chan_count = 0, + const nested_vector2& first_chan_weight_remap = {}, + std::size_t first_remapped_chan_count = 0, + const std::vector& second_chan_weight_remap = {}, + std::size_t second_remapped_chan_count = 0, bool madnis_training = false, bool drop_cuts_and_rescale = false, bool partial_weights = false, @@ -60,7 +63,10 @@ class Integrand : public FunctionGenerator { const nested_vector2& active_flavors = {}, const std::vector& flavor_remap = {}, const std::vector& flavor_factors = {}, - const std::vector& flavor_mirror = {} + const std::vector& flavor_mirror = {}, + const std::vector& flavor_diff_xs_indices = {}, + const std::vector& flavor_subproc_indices = {}, + const std::vector& flavor_per_subproc_remap = {} ); std::size_t particle_count() const { return _mapping.particle_count(); } bool madnis_training() const { return _madnis_training; } @@ -86,10 +92,10 @@ class Integrand : public FunctionGenerator { } } const PhaseSpaceMapping& mapping() const { return _mapping; } - const DifferentialCrossSection& diff_xs() const { return _diff_xs; } + const std::vector& diff_xs() const { return _diff_xs; } const AdaptiveMapping& adaptive_map() const { return _adaptive_map; } - const AdaptiveDiscrete& discrete_before() const { return _discrete_before; } - const AdaptiveDiscrete& discrete_after() const { return _discrete_after; } + const AdaptiveDiscrete& discrete_sym() const { return _discrete_sym; } + const AdaptiveDiscrete& discrete_flavor() const { return _discrete_flavor; } const std::optional& energy_scale() const { return _energy_scale; } const std::optional& prop_chan_weights() const { return _prop_chan_weights; @@ -113,10 +119,11 @@ class Integrand : public FunctionGenerator { build_common_part(FunctionBuilder& fb, const NamedVector& channel_out) const; PhaseSpaceMapping _mapping; - DifferentialCrossSection _diff_xs; + std::vector _diff_xs; AdaptiveMapping _adaptive_map; - AdaptiveDiscrete _discrete_before; - AdaptiveDiscrete _discrete_after; + AdaptiveDiscrete _discrete_sym; + AdaptiveDiscrete _discrete_flavor; + nested_vector2 _pid_options; std::array, 2> _pdfs; std::array, 2> _pdf_indices; std::optional _running_coupling; @@ -124,8 +131,10 @@ class Integrand : public FunctionGenerator { std::optional _prop_chan_weights; std::optional _subchan_weights; std::optional _chan_weight_net; - std::vector _chan_weight_remap; - me_int_t _remapped_chan_count; + nested_vector2 _first_chan_weight_remap; + me_int_t _first_remapped_chan_count; + std::vector _second_chan_weight_remap; + me_int_t _second_remapped_chan_count; bool _madnis_training; bool _drop_cuts_and_rescale; bool _partial_weights; @@ -139,6 +148,9 @@ class Integrand : public FunctionGenerator { std::vector _flavor_mirror; bool _has_mirror; NamedVector _channel_part_ret_types; + std::vector _flavor_diff_xs_indices; + std::vector _flavor_subproc_indices; + std::vector _flavor_per_subproc_remap; friend class IntegrandProbability; friend class IntegrandChannelPart; @@ -209,8 +221,8 @@ class IntegrandProbability : public FunctionGenerator { ) const override; Integrand::AdaptiveMapping _adaptive_map; - Integrand::AdaptiveDiscrete _discrete_before; - Integrand::AdaptiveDiscrete _discrete_after; + Integrand::AdaptiveDiscrete _discrete_sym; + Integrand::AdaptiveDiscrete _discrete_flavor; std::size_t _permutation_count; std::size_t _flavor_count; bool _has_pdf_prior; diff --git a/madspace/instruction_set.yaml b/madspace/instruction_set.yaml index cc8055381..e29401b4a 100644 --- a/madspace/instruction_set.yaml +++ b/madspace/instruction_set.yaml @@ -149,6 +149,26 @@ accept_norm: class: AcceptNormInstruction custom_op: True +batch_split_by_index: + inputs: + - name: indices + type: [int] + desc: + - name: count + type: [int, single] + desc: + outputs: any + class: BatchSplitByIndexInstruction + custom_op: True + +batch_merge_by_index: + inputs: any + outputs: + - name: output + desc: + class: BatchMergeByIndexInstruction + custom_op: True + --- title: Math diff --git a/madspace/src/compgraphs/instruction.cpp b/madspace/src/compgraphs/instruction.cpp index c5f6d88e9..b05d0486a 100644 --- a/madspace/src/compgraphs/instruction.cpp +++ b/madspace/src/compgraphs/instruction.cpp @@ -485,6 +485,90 @@ TypeVec BatchSplitInstruction::signature(const ValueVec& args) const { return out_types; } +TypeVec BatchSplitByIndexInstruction::signature(const ValueVec& args) const { + check_arg_count(args, 2); + auto& indices_type = args.at(0).type; + if (indices_type.dtype != DataType::dt_int || indices_type.shape.size() != 0) { + throw std::invalid_argument( + std::format("{}, argument 1: expected batch of integers", name()) + ); + } + if (indices_type.batch_size == BatchSize::one) { + throw std::invalid_argument( + std::format("{}, argument 1: must have batch dimension", name()) + ); + } + int size = int_literal_arg(args, 1); + TypeVec output_types; + auto last_batch_size = indices_type.batch_size; + for (int i = 0; i < size; ++i) { + auto batch_size = i == size - 1 ? last_batch_size : BatchSize(); + output_types.push_back({DataType::dt_int, batch_size, {}}); + last_batch_size = last_batch_size - batch_size; + } + return output_types; +} + +TypeVec BatchMergeByIndexInstruction::signature(const ValueVec& args) const { + if (args.size() < 2 || args.size() % 2 != 0) { + throw std::invalid_argument( + std::format( + "{} has to be called with an even, positive number of arguments " + "(alternating values and indices)", + name() + ) + ); + } + + auto& type = args.at(0).type; + auto batch_size = BatchSize::zero; + for (std::size_t i = 0; i < args.size(); i += 2) { + auto& values_type = args.at(i).type; + auto& indices_type = args.at(i + 1).type; + if (values_type.dtype == DataType::batch_sizes) { + throw std::invalid_argument( + std::format( + "{}, argument {}: batch size list not accepted as argument", + name(), + i + 1 + ) + ); + } + if (values_type.batch_size == BatchSize::one) { + throw std::invalid_argument( + std::format("{}, argument {}: must have batch dimension", name(), i + 1) + ); + } + if (values_type.dtype != type.dtype || values_type.shape != type.shape) { + throw std::invalid_argument( + std::format( + "{}: all values arguments must have the same shape and dtype", + name() + ) + ); + } + if (indices_type.dtype != DataType::dt_int || indices_type.shape.size() != 0) { + throw std::invalid_argument( + std::format( + "{}, argument {}: expected batch of integers", name(), i + 2 + ) + ); + } + if (indices_type.batch_size != values_type.batch_size) { + throw std::invalid_argument( + std::format( + "{}, argument {}: must have the same batch size as argument {}", + name(), + i + 2, + i + 1 + ) + ); + } + batch_size = batch_size + values_type.batch_size; + } + return {{type.dtype, batch_size, type.shape}}; +} + TypeVec CatInstruction::signature(const ValueVec& args) const { if (args.size() == 0) { throw std::invalid_argument("cat has to be called with at least one argument"); diff --git a/madspace/src/compgraphs/instruction_set_mixin.inc b/madspace/src/compgraphs/instruction_set_mixin.inc index a2dbe6cd8..1e17a1766 100644 --- a/madspace/src/compgraphs/instruction_set_mixin.inc +++ b/madspace/src/compgraphs/instruction_set_mixin.inc @@ -29,147 +29,149 @@ InstructionOwner instructions[] { InstructionOwner(new SqueezeInstruction(11, true)), InstructionOwner(new UnsqueezeInstruction(12, true)), InstructionOwner(new AcceptNormInstruction(13, true)), - 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}}), + InstructionOwner(new BatchSplitByIndexInstruction(14, true)), + InstructionOwner(new BatchMergeByIndexInstruction(15, true)), + mi("add", 16, 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", 17, true, {{DataType::dt_int, false, {std::monostate{}}, false}, {DataType::dt_int, false, {std::monostate{}}, false}}, {{DataType::dt_int, false, {std::monostate{}}, false}}), + mi("sub", 18, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("mul", 19, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("div", 20, 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", 21, true, {{DataType::dt_float, false, {std::monostate{}, "n"}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("reduce_sum_vector", 22, true, {{DataType::dt_float, false, {std::monostate{}, "n", "m"}, false}}, {{DataType::dt_float, false, {std::monostate{}, "m"}, false}}), + mi("batch_reduce_mean", 23, true, {{DataType::dt_float, false, {}, false}}, {{DataType::dt_float, true, {}, false}}), + mi("batch_reduce_mean_keepdim", 24, true, {{DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("reduce_product", 25, true, {{DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("sqrt", 26, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("square", 27, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("min", 28, true, {{DataType::dt_float, false, {std::monostate{}}, false}, {DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("max", 29, 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", 30, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("obs_e", 31, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_px", 32, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_py", 33, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_pz", 34, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_mass", 35, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_pt", 36, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_p_mag", 37, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_phi", 38, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_theta", 39, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_y", 40, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_y_abs", 41, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_eta", 42, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_eta_abs", 43, true, {{DataType::dt_float, false, {std::monostate{}, 4}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("obs_delta_eta", 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_phi", 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("obs_delta_r", 46, 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", 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("boost_beam_inverse", 48, 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", 49, true, {{DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {4}, false}, {DataType::dt_float, false, {4}, false}}), + mi("r_to_x1x2", 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}}), + mi("x1x2_to_r", 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}}), + mi("diff_cross_section", 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, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("two_body_decay_com", 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, {}, false}}), + mi("two_body_decay_com_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, {}, false}}), + mi("two_body_decay", 55, 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", 56, 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", 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_com_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_two_particle_scattering", 59, 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", 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, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("two_to_three_particle_scattering", 61, 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", 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, {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", 63, 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", 64, 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", 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, {}, false}}), + mi("three_body_decay_com_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, {}, false}}), + mi("three_body_decay", 67, 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", 68, 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", 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}}), + mi("t_inv_value_and_min_max", 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}}), + mi("t_inv_min_max_cut", 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("t_inv_value_and_min_max_cut", 72, 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", 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}}), + mi("t1_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}}), + mi("t2_inv_min_max_doublet", 75, 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", 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}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("s23_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, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}), + mi("s23_value_and_min_max", 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}}), + mi("s23_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, {}, 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", 80, 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", 81, 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", 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_float, false, {"c"}, false}}), + mi("subchannel_weights", 83, 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", 84, 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", 85, 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", 86, true, {{DataType::dt_float, false, {"n", 4}, false}, {DataType::dt_int, false, {}, false}}, {{DataType::dt_float, false, {"n", 4}, false}}), + mi("momenta_to_x1x2", 87, 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", 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("uniform_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}}), + mi("breit_wigner_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}, {DataType::dt_float, false, {}, false}}), + mi("breit_wigner_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}, {DataType::dt_float, false, {}, false}}), + mi("stable_invariant", 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_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}}), + mi("stable_invariant_nu", 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("stable_invariant_nu_inverse", 95, 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", 96, 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", 97, 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", 98, 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", 99, 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", 100, 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", 101, 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", 102, 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", 103, true, {{DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}, {DataType::dt_float, false, {}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("cut_all", 104, 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", 105, 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", 106, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("scale_transverse_mass", 107, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("scale_half_transverse_mass", 108, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("scale_partonic_energy", 109, true, {{DataType::dt_float, false, {"n", 4}, false}}, {{DataType::dt_float, false, {}, false}}), + mi("chili_forward", 110, 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", 111, 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(112, true)), + mi("collect_channel_weights", 113, 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", 114, 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", 115, 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", 116, 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", 117, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("leaky_relu", 118, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("elu", 119, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("gelu", 120, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("sigmoid", 121, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("softplus", 122, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + InstructionOwner(new RqsReshapeInstruction(123, true)), + mi("rqs_find_bin", 124, 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", 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("rqs_inverse", 126, 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", 127, true, {{DataType::dt_float, false, {std::monostate{}}, false}}, {{DataType::dt_float, false, {std::monostate{}}, false}}), + mi("softmax_prior", 128, true, {{DataType::dt_float, false, {"n"}, false}, {DataType::dt_float, false, {"n"}, false}}, {{DataType::dt_float, false, {"n"}, false}}), + mi("sample_discrete", 129, 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", 130, 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", 131, 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", 132, 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", 133, 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", 134, 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", 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.cpp b/madspace/src/cpu/runtime.cpp index 4c7f05474..fc206e661 100644 --- a/madspace/src/cpu/runtime.cpp +++ b/madspace/src/cpu/runtime.cpp @@ -379,17 +379,10 @@ void batch_scatter_impl( } template -void op_batch_scatter( - const CpuRuntime::Instruction& instruction, TensorVec& locals, const D& device +void batch_scatter_dispatch( + Tensor& indices, Tensor& source, Tensor& output, const D& device ) { - auto& indices = locals[instruction.input_indices[0]]; - auto& target = locals[instruction.input_indices[1]]; - auto& source = locals[instruction.input_indices[2]]; - - auto& output = locals[instruction.output_indices[0]]; - output = target.copy(device); - device.sync_barrier(); - switch (target.shape().size()) { + switch (output.shape().size()) { case 1: batch_scatter_impl<1>(indices, source, output, device); break; @@ -407,6 +400,73 @@ void op_batch_scatter( } } +template +void op_batch_scatter( + const CpuRuntime::Instruction& instruction, TensorVec& locals, const D& device +) { + auto& indices = locals[instruction.input_indices[0]]; + auto& target = locals[instruction.input_indices[1]]; + auto& source = locals[instruction.input_indices[2]]; + + auto& output = locals[instruction.output_indices[0]]; + output = target.copy(device); + device.sync_barrier(); + batch_scatter_dispatch(indices, source, output, device); +} + +template +void op_batch_split_by_index( + const CpuRuntime::Instruction& instruction, TensorVec& locals, const D& device +) { + auto& indices = locals[instruction.input_indices[0]]; + std::size_t count = locals[instruction.input_indices[1]].index_value(); + auto indices_view_flat = indices.flat_view(0); + device.submit([indices_view_flat, count, &locals, &instruction, &device]() mutable { + std::vector sizes(count); + TensorView indices_view(indices_view_flat); + for (std::size_t i = 0; i < indices_view.size(); ++i) { + sizes[indices_view[i]] += 1; + } + std::vector> views; + views.reserve(count); + for (auto [size, output_index] : zip(sizes, instruction.output_indices)) { + auto& output = locals[output_index]; + output = Tensor(DataType::dt_int, {size}, device); + views.push_back(output.view()); + } + std::fill(sizes.begin(), sizes.end(), 0); + for (std::size_t i = 0; i < indices_view.size(); ++i) { + me_int_t index = indices_view[i]; + std::size_t& size = sizes[index]; + views[index][size] = i; + ++size; + } + }); +} + +template +void op_batch_merge_by_index( + const CpuRuntime::Instruction& instruction, TensorVec& locals, const D& device +) { + std::size_t batch_size = 0; + for (std::size_t i = 0; i < instruction.input_indices.size(); i += 2) { + batch_size += locals[instruction.input_indices[i]].size(0); + } + auto& arg0 = locals[instruction.input_indices[0]]; + auto& output = locals[instruction.output_indices[0]]; + Sizes shape = arg0.shape(); + shape[0] = batch_size; + output = Tensor(arg0.dtype(), shape, device); + for (std::size_t i = 0; i < instruction.input_indices.size(); i += 2) { + batch_scatter_dispatch( + locals[instruction.input_indices[i + 1]], + locals[instruction.input_indices[i]], + output, + device + ); + } +} + template void batch_reduce_mean_impl( const CpuRuntime::Instruction& instruction, diff --git a/madspace/src/cpu/runtime_backward_mixin.inc b/madspace/src/cpu/runtime_backward_mixin.inc index 8a75bf498..dc830f727 100644 --- a/madspace/src/cpu/runtime_backward_mixin.inc +++ b/madspace/src/cpu/runtime_backward_mixin.inc @@ -25,93 +25,93 @@ case 11: case 12: backward_op_unsqueeze(instr, locals, local_grads, device); break; -case 14: +case 16: backward_batch_foreach, backward_kernel_add, 3, 2, DeviceType>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 16: +case 18: backward_batch_foreach, backward_kernel_sub, 3, 2, DeviceType>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 17: +case 19: backward_batch_foreach, backward_kernel_mul, 3, 2, DeviceType>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 18: +case 20: backward_batch_foreach, backward_kernel_div, 3, 2, DeviceType>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 21: +case 23: backward_op_batch_reduce_mean(instr, locals, local_grads, device); break; -case 22: +case 24: backward_op_batch_reduce_mean_keepdim(instr, locals, local_grads, device); break; -case 23: +case 25: 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 26: backward_batch_foreach, backward_kernel_sqrt, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 25: +case 27: backward_batch_foreach, backward_kernel_square, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 114: +case 116: backward_op_matmul(instr, locals, local_grads, device); break; -case 115: +case 117: backward_batch_foreach, backward_kernel_relu, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 116: +case 118: backward_batch_foreach, backward_kernel_leaky_relu, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 117: +case 119: backward_batch_foreach, backward_kernel_elu, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 118: +case 120: backward_batch_foreach, backward_kernel_gelu, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 119: +case 121: backward_batch_foreach, backward_kernel_sigmoid, 2, 1, DeviceType>, 1, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 120: +case 122: backward_batch_foreach, backward_kernel_softplus, 2, 1, DeviceType>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 121: +case 123: backward_op_rqs_reshape(instr, locals, local_grads, device); break; -case 122: +case 124: 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 125: 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 126: 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 127: backward_batch_foreach, backward_kernel_softmax, 2, 1, DeviceType>, 1, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 126: +case 128: 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 132: 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..68d1f9081 100644 --- a/madspace/src/cpu/runtime_mixin.inc +++ b/madspace/src/cpu/runtime_mixin.inc @@ -44,431 +44,437 @@ case 13: op_accept_norm(instr, locals, device); break; case 14: - batch_foreach, kernel_add, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + op_batch_split_by_index(instr, locals, device); break; case 15: - batch_foreach, kernel_add_int, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + op_batch_merge_by_index(instr, locals, device); break; case 16: - batch_foreach, kernel_sub, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_add, 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_add_int, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 18: - batch_foreach, kernel_div, 2, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_sub, 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_mul, 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_div, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 21: - op_batch_reduce_mean(instr, locals, device); + batch_foreach, kernel_reduce_sum, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 22: - op_batch_reduce_mean_keepdim(instr, locals, device); + batch_foreach, kernel_reduce_sum_vector, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 23: - batch_foreach, kernel_reduce_product, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + op_batch_reduce_mean(instr, locals, device); break; case 24: - batch_foreach, kernel_sqrt, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + op_batch_reduce_mean_keepdim(instr, locals, device); break; case 25: - batch_foreach, kernel_square, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_reduce_product, 1, 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_sqrt, 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_square, 1, 1, DeviceType>, 1, 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_min, 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_max, 2, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 30: - batch_foreach, kernel_obs_px, 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 31: - batch_foreach, kernel_obs_py, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_obs_e, 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_px, 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_py, 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_pz, 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_mass, 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_pt, 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_p_mag, 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_phi, 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_theta, 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, 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_y_abs, 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, 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_eta_abs, 1, 1, DeviceType>, 1, 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_eta, 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_phi, 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_obs_delta_r, 2, 1, DeviceType>, 2, 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, 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_boost_beam_inverse, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 49: - batch_foreach, kernel_x1x2_to_r, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); + batch_foreach, kernel_com_p_in, 1, 2, 1, DeviceType>, 1, 2>(instr, locals, device); break; case 50: - batch_foreach, kernel_diff_cross_section, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); + batch_foreach, kernel_r_to_x1x2, 3, 3, 1, DeviceType>, 3, 3>(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_x1x2_to_r, 3, 2, 1, DeviceType>, 3, 2>(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_diff_cross_section, 6, 1, 1, DeviceType>, 6, 1>(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, 5, 3, 1, DeviceType>, 5, 3>(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_com_inverse, 2, 6, 1, DeviceType>, 2, 6>(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, 6, 3, 1, DeviceType>, 6, 3>(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_body_decay_inverse, 2, 7, 1, DeviceType>, 2, 7>(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, 6, 3, 1, DeviceType>, 6, 3>(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_com_inverse, 4, 4, 1, DeviceType>, 4, 4>(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, 6, 3, 1, DeviceType>, 6, 3>(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_two_particle_scattering_inverse, 4, 4, 1, DeviceType>, 4, 4>(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, 8, 3, 1, DeviceType>, 8, 3>(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_two_to_three_particle_scattering_inverse, 7, 4, 1, DeviceType>, 7, 4>(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, 6, 3, 1, DeviceType>, 6, 3>(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_double_t_scattering_inverse, 4, 2, 1, DeviceType>, 4, 2>(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, 9, 4, 1, DeviceType>, 9, 4>(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_com_inverse, 3, 10, 1, DeviceType>, 3, 10>(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, 10, 4, 1, DeviceType>, 10, 4>(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_three_body_decay_inverse, 3, 11, 1, DeviceType>, 3, 11>(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_min_max, 4, 2, 1, DeviceType>, 4, 2>(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_value_and_min_max, 4, 3, 1, DeviceType>, 4, 3>(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_min_max_cut, 6, 2, 1, DeviceType>, 6, 2>(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_t_inv_value_and_min_max_cut, 6, 3, 1, DeviceType>, 6, 3>(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_min_max_doublet, 6, 2, 1, DeviceType>, 6, 2>(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_t1_inv_value_and_min_max_doublet, 7, 3, 1, DeviceType>, 7, 3>(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_min_max_doublet, 7, 2, 1, DeviceType>, 7, 2>(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_t2_inv_value_and_min_max_doublet, 8, 3, 1, DeviceType>, 8, 3>(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_min_max, 6, 2, 1, DeviceType>, 6, 2>(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_value_and_min_max, 6, 3, 1, DeviceType>, 6, 3>(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_min_max_cut, 10, 2, 1, DeviceType>, 10, 2>(instr, locals, device); break; case 80: - batch_foreach, kernel_sde2_channel_weights, 4, 1, 1, DeviceType>, 4, 1>(instr, locals, device); + batch_foreach, kernel_s23_value_and_min_max_cut, 10, 3, 1, DeviceType>, 10, 3>(instr, locals, device); break; case 81: - batch_foreach, kernel_subchannel_weights, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); + batch_foreach, kernel_invariants_from_momenta, 2, 1, 1, DeviceType>, 2, 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_sde2_channel_weights, 4, 1, 1, DeviceType>, 4, 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_subchannel_weights, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); break; case 84: - batch_foreach, kernel_mirror_momenta, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + batch_foreach, kernel_apply_subchannel_weights, 4, 1, 1, DeviceType>, 4, 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_pt_eta_phi_x, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 86: - batch_foreach, kernel_uniform_invariant, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); + batch_foreach, kernel_mirror_momenta, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); break; case 87: - batch_foreach, kernel_uniform_invariant_inverse, 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 88: - batch_foreach, kernel_breit_wigner_invariant, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); + batch_foreach, kernel_uniform_invariant, 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_uniform_invariant_inverse, 3, 2, 1, DeviceType>, 3, 2>(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, 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_breit_wigner_invariant_inverse, 5, 2, 1, DeviceType>, 5, 2>(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, 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_inverse, 4, 2, 1, DeviceType>, 4, 2>(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, 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_stable_invariant_nu_inverse, 5, 2, 1, DeviceType>, 5, 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, 3, 2, 1, DeviceType>, 3, 2>(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_inverse, 2, 3, 1, DeviceType>, 2, 3>(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_massless_com, 2, 2, 1, DeviceType>, 2, 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, 4, 2, 1, DeviceType>, 4, 2>(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_inverse, 3, 3, 1, DeviceType>, 3, 3>(instr, locals, device); break; case 101: - batch_foreach, kernel_cut_one, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + batch_foreach, kernel_fast_rambo_massive_com, 3, 2, 1, DeviceType>, 3, 2>(instr, locals, device); break; case 102: - batch_foreach, kernel_cut_all, 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 103: - batch_foreach, kernel_cut_any, 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 104: - batch_foreach, kernel_scale_transverse_energy, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_cut_all, 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_cut_any, 3, 1, 1, DeviceType>, 3, 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_energy, 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_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_half_transverse_mass, 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_scale_partonic_energy, 1, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 110: - op_matrix_element(instr, locals, device); + batch_foreach, kernel_chili_forward, 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); + batch_foreach, kernel_chili_inverse, 5, 2, 1, DeviceType>, 5, 2>(instr, locals, device); break; case 112: - batch_foreach, kernel_interpolate_pdf, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); + op_matrix_element(instr, locals, device); break; case 113: - batch_foreach, kernel_interpolate_alpha_s, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); + batch_foreach, kernel_collect_channel_weights, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 114: - op_matmul(instr, locals, device); + batch_foreach, kernel_interpolate_pdf, 6, 1, 1, DeviceType>, 6, 1>(instr, locals, device); break; case 115: - batch_foreach, kernel_relu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_interpolate_alpha_s, 3, 1, 1, DeviceType>, 3, 1>(instr, locals, device); break; case 116: - batch_foreach, kernel_leaky_relu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + op_matmul(instr, locals, device); break; case 117: - batch_foreach, kernel_elu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_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_leaky_relu, 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_elu, 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_gelu, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 121: - op_rqs_reshape(instr, locals, device); + batch_foreach, kernel_sigmoid, 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); + batch_foreach, kernel_softplus, 1, 1, DeviceType>, 1, 1>(instr, locals, device); break; case 123: - batch_foreach, kernel_rqs_forward, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); + op_rqs_reshape(instr, locals, device); break; case 124: - batch_foreach, kernel_rqs_inverse, 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 125: - batch_foreach, kernel_softmax, 1, 1, DeviceType>, 1, 1>(instr, locals, device); + batch_foreach, kernel_rqs_forward, 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_rqs_inverse, 2, 2, 2, DeviceType>, 2, 2>(instr, locals, device); break; case 127: - batch_foreach, kernel_sample_discrete, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); + batch_foreach, kernel_softmax, 1, 1, DeviceType>, 1, 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_softmax_prior, 2, 1, 1, DeviceType>, 2, 1>(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, 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_inverse, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); break; case 131: - op_discrete_histogram(instr, locals, device); + batch_foreach, kernel_sample_discrete_probs, 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); + batch_foreach, kernel_sample_discrete_probs_inverse, 2, 2, 1, DeviceType>, 2, 2>(instr, locals, device); break; case 133: - batch_foreach, kernel_gather, 2, 1, 1, DeviceType>, 2, 1>(instr, locals, device); + op_discrete_histogram(instr, locals, device); break; case 134: - batch_foreach, kernel_gather_int, 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 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/driver/channel_generator.cpp b/madspace/src/driver/channel_generator.cpp index c86da5107..3f1601bfc 100644 --- a/madspace/src/driver/channel_generator.cpp +++ b/madspace/src/driver/channel_generator.cpp @@ -16,6 +16,9 @@ int event_extra_flags(const std::unordered_map& index_ if (index_map.contains("partial_weight_product")) { flags |= EventRecord::f_partial_weights; } + if (index_map.contains("subprocess_index")) { + flags |= EventRecord::f_subproc_index; + } return flags; } @@ -120,8 +123,8 @@ ChannelEventGenerator::ChannelEventGenerator( }, [](auto sampler) {} }; - std::visit(add_names, integrand.discrete_before()); - std::visit(add_names, integrand.discrete_after()); + std::visit(add_names, integrand.discrete_sym()); + std::visit(add_names, integrand.discrete_flavor()); std::optional discrete_optimizer; RuntimePtr discrete_histogram = nullptr; if (prob_names.size() > 0) { @@ -271,6 +274,11 @@ void ChannelEventGenerator::init_field_indices() { } else { _field_indices.partial_weight_product = -1; } + if (index_map.contains("subprocess_index")) { + _field_indices.subprocess_index = index_map.at("subprocess_index"); + } else { + _field_indices.subprocess_index = -1; + } _field_indices.random = index_map.at("random"); _field_indices.rest = _field_indices.random + 1; } @@ -643,6 +651,15 @@ void ChannelEventGenerator::write_events( } } + if (_field_indices.subprocess_index != -1) { + auto subproc_view = + unweighted_events.at(_field_indices.subprocess_index).view(); + for (std::size_t i = 0; i < subproc_view.size(); ++i) { + auto event = event_buffer.event(i); + event.subprocess_index() = subproc_view[i]; + } + } + _event_file.write(event_buffer); _weight_file.write(weight_buffer); _status.count_unweighted += diff --git a/madspace/src/driver/event_generator.cpp b/madspace/src/driver/event_generator.cpp index af8e3a59b..2a6ec09d0 100644 --- a/madspace/src/driver/event_generator.cpp +++ b/madspace/src/driver/event_generator.cpp @@ -606,6 +606,8 @@ void EventGenerator::read_and_combine( bool has_beam2 = _channels.at(0)->event_layout_extra_flags() & EventRecord::f_beam2; bool has_partial = _channels.at(0)->event_layout_extra_flags() & EventRecord::f_partial_weights; + bool has_subproc_index = + _channels.at(0)->event_layout_extra_flags() & EventRecord::f_subproc_index; std::random_device rand_device; std::mt19937 rand_gen(rand_device()); @@ -642,7 +644,9 @@ void EventGenerator::read_and_combine( auto event_in = sampled_chan->event_buffer.event(sampled_chan->buffer_index); auto event_out = buffer.event(event_index); event_out.weight() = std::max(1., weight / channel->max_weight()) * norm_factor; - event_out.subprocess_index() = channel->status().subprocess; + event_out.subprocess_index() = has_subproc_index + ? event_in.subprocess_index().value() + : static_cast(channel->status().subprocess); event_out.diagram_index() = event_in.diagram_index(); event_out.color_index() = event_in.color_index(); event_out.flavor_index() = event_in.flavor_index(); diff --git a/madspace/src/gpu/runtime.cu b/madspace/src/gpu/runtime.cu index 87ff0b0be..a21fd88c1 100644 --- a/madspace/src/gpu/runtime.cu +++ b/madspace/src/gpu/runtime.cu @@ -13,6 +13,7 @@ #include #include #include +#include #include #include "../kernels/kernels.hpp" @@ -506,18 +507,10 @@ void batch_scatter_impl( } } -void op_batch_scatter( - const GpuRuntime::Instruction& instruction, - TensorVec& locals, - const AsyncGpuDevice& device +void batch_scatter_dispatch( + Tensor& indices, Tensor& source, Tensor& output, const AsyncGpuDevice& device ) { - auto& indices = locals[instruction.input_indices[0]]; - auto& target = locals[instruction.input_indices[1]]; - auto& source = locals[instruction.input_indices[2]]; - - auto& output = locals[instruction.output_indices[0]]; - output = target.copy(device, instruction.output_alloc_hints[0]); - switch (target.shape().size()) { + switch (output.shape().size()) { case 1: batch_scatter_impl<1>(indices, source, output, device); break; @@ -535,6 +528,170 @@ void op_batch_scatter( } } +void op_batch_scatter( + const GpuRuntime::Instruction& instruction, + TensorVec& locals, + const AsyncGpuDevice& device +) { + auto& indices = locals[instruction.input_indices[0]]; + auto& target = locals[instruction.input_indices[1]]; + auto& source = locals[instruction.input_indices[2]]; + + auto& output = locals[instruction.output_indices[0]]; + output = target.copy(device, instruction.output_alloc_hints[0]); + batch_scatter_dispatch(indices, source, output, device); +} + +void op_batch_split_by_index( + const GpuRuntime::Instruction& instruction, + TensorVec& locals, + const AsyncGpuDevice& device +) { + auto indices = + locals[instruction.input_indices[0]].contiguous(device, AllocHint::temporary); + std::size_t count = locals[instruction.input_indices[1]].index_value(); + std::size_t batch_size = indices.size(0); + + if (batch_size == 0) { + for (std::size_t k = 0; k < count; ++k) { + locals[instruction.output_indices[k]] = Tensor( + DataType::dt_int, {0}, device, instruction.output_alloc_hints[k] + ); + } + return; + } + + // per-bucket counts, needed on the host to size and split the outputs below + Tensor sizes(DataType::dt_int, {count}, device, AllocHint::temporary); + { + std::size_t temp_storage_bytes = 0; + check_error( + cub::DeviceHistogram::HistogramEven( + nullptr, + temp_storage_bytes, + static_cast(indices.data()), + static_cast(sizes.data()), + static_cast(count) + 1, + 0, + static_cast(count), + batch_size, + device.stream() + ) + ); + Tensor temp( + DataType::dt_float, + {(temp_storage_bytes + 7) / 8}, + device, + AllocHint::temporary + ); + check_error( + cub::DeviceHistogram::HistogramEven( + temp.data(), + temp_storage_bytes, + static_cast(indices.data()), + static_cast(sizes.data()), + static_cast(count) + 1, + 0, + static_cast(count), + batch_size, + device.stream() + ) + ); + temp.reset(device); + } + + // stable sort of the original positions by bucket index: afterwards each + // bucket's positions are contiguous, in their original relative order + Tensor values_in(DataType::dt_int, {batch_size}, device, AllocHint::temporary); + auto values_in_ptr = + thrust::device_pointer_cast(static_cast(values_in.data())); + thrust::sequence( + thrust_par.on(device.stream()), values_in_ptr, values_in_ptr + batch_size + ); + Tensor keys_out(DataType::dt_int, {batch_size}, device, AllocHint::temporary); + Tensor values_out( + DataType::dt_int, {batch_size}, device, instruction.output_alloc_hints[0] + ); + { + std::size_t temp_storage_bytes = 0; + check_error( + cub::DeviceRadixSort::SortPairs( + nullptr, + temp_storage_bytes, + static_cast(indices.data()), + static_cast(keys_out.data()), + static_cast(values_in.data()), + static_cast(values_out.data()), + batch_size, + 0, + static_cast(sizeof(me_int_t) * 8), + device.stream() + ) + ); + Tensor temp( + DataType::dt_float, + {(temp_storage_bytes + 7) / 8}, + device, + AllocHint::temporary + ); + check_error( + cub::DeviceRadixSort::SortPairs( + temp.data(), + temp_storage_bytes, + static_cast(indices.data()), + static_cast(keys_out.data()), + static_cast(values_in.data()), + static_cast(values_out.data()), + batch_size, + 0, + static_cast(sizeof(me_int_t) * 8), + device.stream() + ) + ); + temp.reset(device); + } + values_in.reset(device); + keys_out.reset(device); + indices.reset(device); + + Tensor sizes_cpu = sizes.cpu(device); + check_error(gpuStreamSynchronize(device.stream())); + auto sizes_view = sizes_cpu.view(); + std::size_t offset = 0; + for (std::size_t k = 0; k < count; ++k) { + std::size_t bucket_size = sizes_view[k]; + locals[instruction.output_indices[k]] = + values_out.slice(0, offset, offset + bucket_size); + offset += bucket_size; + } + values_out.reset(device); + sizes.reset(device); +} + +void op_batch_merge_by_index( + const GpuRuntime::Instruction& instruction, + TensorVec& locals, + const AsyncGpuDevice& device +) { + std::size_t batch_size = 0; + for (std::size_t i = 0; i < instruction.input_indices.size(); i += 2) { + batch_size += locals[instruction.input_indices[i]].size(0); + } + auto& arg0 = locals[instruction.input_indices[0]]; + auto& output = locals[instruction.output_indices[0]]; + Sizes shape = arg0.shape(); + shape[0] = batch_size; + output = Tensor(arg0.dtype(), shape, device, instruction.output_alloc_hints[0]); + for (std::size_t i = 0; i < instruction.input_indices.size(); i += 2) { + batch_scatter_dispatch( + locals[instruction.input_indices[i + 1]], + locals[instruction.input_indices[i]], + output, + device + ); + } +} + __global__ void kernel_div_batch_size( std::size_t batch_size, GpuTensorView input, @@ -559,24 +716,28 @@ void batch_reduce_mean_impl( Tensor out(DataType::dt_float, {1}, device, hint); std::size_t temp_storage_bytes = 0; - cub::DeviceReduce::Sum( - nullptr, - temp_storage_bytes, - static_cast(input.data()), - static_cast(out.data()), - batch_size, - device.stream() + check_error( + cub::DeviceReduce::Sum( + nullptr, + temp_storage_bytes, + static_cast(input.data()), + static_cast(out.data()), + batch_size, + device.stream() + ) ); Tensor temp( DataType::dt_float, {(temp_storage_bytes + 7) / 8}, device, AllocHint::temporary ); - cub::DeviceReduce::Sum( - temp.data(), - temp_storage_bytes, - static_cast(input.data()), - static_cast(out.data()), - batch_size, - device.stream() + check_error( + cub::DeviceReduce::Sum( + temp.data(), + temp_storage_bytes, + static_cast(input.data()), + static_cast(out.data()), + batch_size, + device.stream() + ) ); temp.reset(device); input.reset(device); @@ -620,13 +781,15 @@ void batch_reduce_mean_backward_impl( local_grads[instruction.output_indices[0]].contiguous(device); std::size_t batch_size = output_grad.size(0); std::size_t temp_storage_bytes = 0; - cub::DeviceReduce::Sum( - nullptr, - temp_storage_bytes, - static_cast(output_grad.data()), - static_cast(grad.data()), - batch_size, - device.stream() + check_error( + cub::DeviceReduce::Sum( + nullptr, + temp_storage_bytes, + static_cast(output_grad.data()), + static_cast(grad.data()), + batch_size, + device.stream() + ) ); Tensor temp( DataType::dt_float, @@ -634,13 +797,15 @@ void batch_reduce_mean_backward_impl( device, AllocHint::temporary ); - cub::DeviceReduce::Sum( - temp.data(), - temp_storage_bytes, - static_cast(output_grad.data()), - static_cast(grad.data()), - batch_size, - device.stream() + check_error( + cub::DeviceReduce::Sum( + temp.data(), + temp_storage_bytes, + static_cast(output_grad.data()), + static_cast(grad.data()), + batch_size, + device.stream() + ) ); launch_kernel( kernel_div_batch_size, @@ -763,24 +928,28 @@ void op_quantile( Tensor tmp(DataType::dt_float, {batch_size}, device, AllocHint::temporary); tmp.copy_from(input, device); std::size_t temp_storage_bytes; - cub::DeviceMergeSort::SortKeys( - nullptr, - temp_storage_bytes, - static_cast(tmp.data()), - batch_size, - gpu_less_double{}, - device.stream() + check_error( + cub::DeviceMergeSort::SortKeys( + nullptr, + temp_storage_bytes, + static_cast(tmp.data()), + batch_size, + gpu_less_double{}, + device.stream() + ) ); Tensor tmp_sort( DataType::dt_float, {(temp_storage_bytes + 7) / 8}, device, AllocHint::temporary ); - cub::DeviceMergeSort::SortKeys( - tmp_sort.data(), - temp_storage_bytes, - static_cast(tmp.data()), - batch_size, - gpu_less_double{}, - device.stream() + check_error( + cub::DeviceMergeSort::SortKeys( + tmp_sort.data(), + temp_storage_bytes, + static_cast(tmp.data()), + batch_size, + gpu_less_double{}, + device.stream() + ) ); tmp_sort.reset(device); diff --git a/madspace/src/gpu/runtime_backward_mixin.inc b/madspace/src/gpu/runtime_backward_mixin.inc index 452946c95..cd854d10e 100644 --- a/madspace/src/gpu/runtime_backward_mixin.inc +++ b/madspace/src/gpu/runtime_backward_mixin.inc @@ -25,93 +25,93 @@ case 11: case 12: backward_op_unsqueeze(instr, locals, local_grads, device); break; -case 14: +case 16: backward_batch_foreach, 3, 2>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 16: +case 18: backward_batch_foreach, 3, 2>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 17: +case 19: backward_batch_foreach, 3, 2>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 18: +case 20: backward_batch_foreach, 3, 2>, 2, 1, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 21: +case 23: backward_op_batch_reduce_mean(instr, locals, local_grads, device); break; -case 22: +case 24: backward_op_batch_reduce_mean_keepdim(instr, locals, local_grads, device); break; -case 23: +case 25: backward_batch_foreach, 2, 1, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 24: +case 26: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 25: +case 27: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 114: +case 116: backward_op_matmul(instr, locals, local_grads, device); break; -case 115: +case 117: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 116: +case 118: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 117: +case 119: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 118: +case 120: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 119: +case 121: backward_batch_foreach, 2, 1>, 1, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 120: +case 122: backward_batch_foreach, 2, 1>, 1, 1, 1, 0>(instr, locals, local_grads, {0}, {}, device); break; -case 121: +case 123: backward_op_rqs_reshape(instr, locals, local_grads, device); break; -case 122: +case 124: backward_batch_foreach, 5, 4, 2>, 4, 1, 4, 0>(instr, locals, local_grads, {0,1,2,3}, {}, device); break; -case 123: +case 125: backward_batch_foreach, 4, 2, 2>, 2, 2, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 124: +case 126: backward_batch_foreach, 4, 2, 2>, 2, 2, 2, 0>(instr, locals, local_grads, {0,1}, {}, device); break; -case 125: +case 127: backward_batch_foreach, 2, 1>, 1, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 126: +case 128: backward_batch_foreach, 2, 2, 1>, 2, 1, 0, 1>(instr, locals, local_grads, {}, {0}, device); break; -case 130: +case 132: 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..825af8d5a 100644 --- a/madspace/src/gpu/runtime_mixin.inc +++ b/madspace/src/gpu/runtime_mixin.inc @@ -44,431 +44,437 @@ case 13: op_accept_norm(instr, locals, device); break; case 14: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + op_batch_split_by_index(instr, locals, device); break; case 15: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + op_batch_merge_by_index(instr, locals, device); break; case 16: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1>, 2, 1>(instr, locals, device); break; case 17: - batch_foreach, 2, 1>, 2, 1>(instr, locals, device); + batch_foreach, 2, 1>, 2, 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, 2, 1>, 2, 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); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 23: - batch_foreach, 1, 1, 1>, 1, 1>(instr, locals, device); + op_batch_reduce_mean(instr, locals, device); break; case 24: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + op_batch_reduce_mean_keepdim(instr, locals, device); break; case 25: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 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, 1, 1>, 1, 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, 2, 1>, 2, 1>(instr, locals, device); break; case 30: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 1, 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, 1, 1>, 1, 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, 2, 1>, 2, 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, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 49: - batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); + batch_foreach, 1, 2, 1>, 1, 2>(instr, locals, device); break; case 50: - batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); + batch_foreach, 3, 3, 1>, 3, 3>(instr, locals, device); break; case 51: - batch_foreach, 5, 3, 1>, 5, 3>(instr, locals, device); + batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); break; case 52: - batch_foreach, 2, 6, 1>, 2, 6>(instr, locals, device); + batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); break; case 53: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 5, 3, 1>, 5, 3>(instr, locals, device); break; case 54: - batch_foreach, 2, 7, 1>, 2, 7>(instr, locals, device); + batch_foreach, 2, 6, 1>, 2, 6>(instr, locals, device); break; case 55: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 56: - batch_foreach, 4, 4, 1>, 4, 4>(instr, locals, device); + batch_foreach, 2, 7, 1>, 2, 7>(instr, locals, device); break; case 57: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 58: - batch_foreach, 4, 4, 1>, 4, 4>(instr, locals, device); + batch_foreach, 4, 4, 1>, 4, 4>(instr, locals, device); break; case 59: - batch_foreach, 8, 3, 1>, 8, 3>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 60: - batch_foreach, 7, 4, 1>, 7, 4>(instr, locals, device); + batch_foreach, 4, 4, 1>, 4, 4>(instr, locals, device); break; case 61: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 8, 3, 1>, 8, 3>(instr, locals, device); break; case 62: - batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); + batch_foreach, 7, 4, 1>, 7, 4>(instr, locals, device); break; case 63: - batch_foreach, 9, 4, 1>, 9, 4>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 64: - batch_foreach, 3, 10, 1>, 3, 10>(instr, locals, device); + batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); break; case 65: - batch_foreach, 10, 4, 1>, 10, 4>(instr, locals, device); + batch_foreach, 9, 4, 1>, 9, 4>(instr, locals, device); break; case 66: - batch_foreach, 3, 11, 1>, 3, 11>(instr, locals, device); + batch_foreach, 3, 10, 1>, 3, 10>(instr, locals, device); break; case 67: - batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); + batch_foreach, 10, 4, 1>, 10, 4>(instr, locals, device); break; case 68: - batch_foreach, 4, 3, 1>, 4, 3>(instr, locals, device); + batch_foreach, 3, 11, 1>, 3, 11>(instr, locals, device); break; case 69: - batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); + batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); break; case 70: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 4, 3, 1>, 4, 3>(instr, locals, device); break; case 71: - batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); + batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); break; case 72: - batch_foreach, 7, 3, 1>, 7, 3>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 73: - batch_foreach, 7, 2, 1>, 7, 2>(instr, locals, device); + batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); break; case 74: - batch_foreach, 8, 3, 1>, 8, 3>(instr, locals, device); + batch_foreach, 7, 3, 1>, 7, 3>(instr, locals, device); break; case 75: - batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); + batch_foreach, 7, 2, 1>, 7, 2>(instr, locals, device); break; case 76: - batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); + batch_foreach, 8, 3, 1>, 8, 3>(instr, locals, device); break; case 77: - batch_foreach, 10, 2, 1>, 10, 2>(instr, locals, device); + batch_foreach, 6, 2, 1>, 6, 2>(instr, locals, device); break; case 78: - batch_foreach, 10, 3, 1>, 10, 3>(instr, locals, device); + batch_foreach, 6, 3, 1>, 6, 3>(instr, locals, device); break; case 79: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 10, 2, 1>, 10, 2>(instr, locals, device); break; case 80: - batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); + batch_foreach, 10, 3, 1>, 10, 3>(instr, locals, device); break; case 81: - batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 82: - batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); + batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); break; case 83: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); break; case 84: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); break; case 85: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 86: - batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); break; case 87: - batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); + batch_foreach, 2, 2, 1>, 2, 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, 3, 2, 1>, 3, 2>(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, 5, 2, 1>, 5, 2>(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, 4, 2, 1>, 4, 2>(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, 5, 2, 1>, 5, 2>(instr, locals, device); break; case 96: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); break; case 97: - batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); + batch_foreach, 2, 3, 1>, 2, 3>(instr, locals, device); break; case 98: - batch_foreach, 3, 3, 1>, 3, 3>(instr, locals, device); + batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); break; case 99: - batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); + batch_foreach, 4, 2, 1>, 4, 2>(instr, locals, device); break; case 100: - batch_foreach, 4, 1, 1>, 4, 1>(instr, locals, device); + batch_foreach, 3, 3, 1>, 3, 3>(instr, locals, device); break; case 101: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 3, 2, 1>, 3, 2>(instr, locals, device); break; case 102: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 4, 1, 1>, 4, 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, 3, 1, 1>, 3, 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, 1, 1, 1>, 1, 1>(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); + batch_foreach, 5, 2, 1>, 5, 2>(instr, locals, device); break; case 112: - batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); + op_matrix_element(instr, locals, device); break; case 113: - batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 114: - op_matmul(instr, locals, device); + batch_foreach, 6, 1, 1>, 6, 1>(instr, locals, device); break; case 115: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + batch_foreach, 3, 1, 1>, 3, 1>(instr, locals, device); break; case 116: - batch_foreach, 1, 1>, 1, 1>(instr, locals, device); + op_matmul(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); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 123: - batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); + op_rqs_reshape(instr, locals, device); break; case 124: - batch_foreach, 2, 2, 2>, 2, 2>(instr, locals, device); + batch_foreach, 4, 1, 2>, 4, 1>(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, 2, 2, 2>, 2, 2>(instr, locals, device); break; case 127: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 1, 1>, 1, 1>(instr, locals, device); break; case 128: - batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); + batch_foreach, 2, 1, 1>, 2, 1>(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); + batch_foreach, 2, 2, 1>, 2, 2>(instr, locals, device); break; case 133: - batch_foreach, 2, 1, 1>, 2, 1>(instr, locals, device); + op_discrete_histogram(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/phasespace/integrand.cpp b/madspace/src/phasespace/integrand.cpp index 2d808ee34..69ea4a3a3 100644 --- a/madspace/src/phasespace/integrand.cpp +++ b/madspace/src/phasespace/integrand.cpp @@ -6,22 +6,48 @@ using namespace madspace; +namespace { + +std::size_t final_channel_count( + const std::vector& diff_xs, + const nested_vector2& first_chan_weight_remap, + std::size_t first_remapped_chan_count, + const std::vector& second_chan_weight_remap, + std::size_t second_remapped_chan_count, + const std::optional& subchan_weights +) { + if (second_chan_weight_remap.size() > 0) { + return second_remapped_chan_count; + } else if (subchan_weights) { + return subchan_weights->channel_count(); + } else if (first_chan_weight_remap.size() > 0) { + return first_remapped_chan_count; + } else { + return diff_xs.at(0).matrix_element().diagram_count(); + } +} + +} // namespace + static const BatchSize acc_batch_size("acc_batch_size"); Integrand::Integrand( const PhaseSpaceMapping& mapping, - const DifferentialCrossSection& diff_xs, + const std::vector& diff_xs, const AdaptiveMapping& adaptive_map, - const AdaptiveDiscrete& discrete_before, - const AdaptiveDiscrete& discrete_after, + const AdaptiveDiscrete& discrete_sym, + const AdaptiveDiscrete& discrete_flavor, + const nested_vector2& pid_options, const std::optional& pdf_grid, const std::optional& running_coupling, const std::optional& energy_scale, const std::optional& prop_chan_weights, const std::optional& subchan_weights, const std::optional& chan_weight_net, - const std::vector& chan_weight_remap, - std::size_t remapped_chan_count, + const nested_vector2& first_chan_weight_remap, + std::size_t first_remapped_chan_count, + const std::vector& second_chan_weight_remap, + std::size_t second_remapped_chan_count, bool madnis_training, bool drop_cuts_and_rescale, bool partial_weights, @@ -29,14 +55,18 @@ Integrand::Integrand( const nested_vector2& active_flavors, const std::vector& flavor_remap, const std::vector& flavor_factors, - const std::vector& flavor_mirror + const std::vector& flavor_mirror, + const std::vector& flavor_diff_xs_indices, + const std::vector& flavor_subproc_indices, + const std::vector& flavor_per_subproc_remap ) : FunctionGenerator( "Integrand", {{"batch_size", Type({batch_size})}}, [&] { NamedVector ret_types; - auto flav_count = diff_xs.pid_options().size(); + auto& diff_xs_first = diff_xs.at(0); + auto flav_count = pid_options.size(); if (madnis_training) { if (std::holds_alternative(adaptive_map)) { throw std::invalid_argument( @@ -50,11 +80,14 @@ Integrand::Integrand( ret_types.push_back("channel_index", batch_int); ret_types.push_back( "channel_weights", - batch_float_array( - chan_weight_remap.size() > 0 ? remapped_chan_count - : subchan_weights ? subchan_weights->channel_count() - : diff_xs.matrix_element().diagram_count() - ) + batch_float_array(final_channel_count( + diff_xs, + first_chan_weight_remap, + first_remapped_chan_count, + second_chan_weight_remap, + second_remapped_chan_count, + subchan_weights + )) ); ret_types.push_back( "cwnet_input", @@ -64,7 +97,7 @@ Integrand::Integrand( ); ret_types.push_back("channel_index_in_group", batch_int); if (flav_count > 1 && - !std::holds_alternative(discrete_after)) { + !std::holds_alternative(discrete_flavor)) { ret_types.push_back("discrete_flavor_index", batch_int); if (pdf_grid && energy_scale) { ret_types.push_back("pdf_prior", batch_float_array(flav_count)); @@ -79,28 +112,31 @@ Integrand::Integrand( ret_types.push_back("helicity_index", batch_int); ret_types.push_back("diagram_index", batch_int); ret_types.push_back("flavor_index", batch_int); + if (flavor_diff_xs_indices.size() > 0) { + ret_types.push_back("subprocess_index", batch_int); + } ret_types.push_back("ren_scale", batch_float); ret_types.push_back("alpha_qcd", batch_float); if (partial_weights) { - if (diff_xs.has_pdf(0)) { + if (diff_xs_first.has_pdf(0)) { ret_types.push_back("x1", batch_float); ret_types.push_back("fact_scale1", batch_float); } - if (diff_xs.has_pdf(1)) { + if (diff_xs_first.has_pdf(1)) { ret_types.push_back("x2", batch_float); ret_types.push_back("fact_scale2", batch_float); } - if (diff_xs.has_pdf(0) || diff_xs.has_pdf(1)) { + if (diff_xs_first.has_pdf(0) || diff_xs_first.has_pdf(1)) { ret_types.push_back("partial_weight_product", batch_float); } } ret_types.push_back("random", batch_float_array(mapping.random_dim())); if (mapping.channel_count() > 1 && - !std::holds_alternative(discrete_before)) { + !std::holds_alternative(discrete_sym)) { ret_types.push_back("channel_index_in_group", batch_int); } if (flav_count > 1 && - !std::holds_alternative(discrete_after)) { + !std::holds_alternative(discrete_flavor)) { ret_types.push_back("discrete_flavor_index", batch_int); } } @@ -110,35 +146,47 @@ Integrand::Integrand( _mapping(mapping), _diff_xs(diff_xs), _adaptive_map(adaptive_map), - _discrete_before(discrete_before), - _discrete_after(discrete_after), + _discrete_sym(discrete_sym), + _discrete_flavor(discrete_flavor), + _pid_options(pid_options), _running_coupling(running_coupling), _energy_scale(energy_scale), _prop_chan_weights(prop_chan_weights), _subchan_weights(subchan_weights), _chan_weight_net(chan_weight_net), - _chan_weight_remap(chan_weight_remap), - _remapped_chan_count(remapped_chan_count), + _first_chan_weight_remap(first_chan_weight_remap), + _first_remapped_chan_count(first_remapped_chan_count), + _second_chan_weight_remap(second_chan_weight_remap), + _second_remapped_chan_count(second_remapped_chan_count), _madnis_training(madnis_training), _drop_cuts_and_rescale(drop_cuts_and_rescale), _partial_weights(partial_weights), _channel_indices(channel_indices.begin(), channel_indices.end()), _random_dim( - mapping.random_dim() + // phasespace - (mapping.channel_count() > 1) + // symmetric channel - (diff_xs.pid_options().size() > 1) + // flavor - // flipped initial state/ + mapping.random_dim() + // phasespace + (mapping.channel_count() > 1) + // symmetric channel + (pid_options.size() > 1) + // flavor + // flipped initial state std::any_of(flavor_mirror.begin(), flavor_mirror.end(), std::identity{}) ), _flavor_remap(flavor_remap.begin(), flavor_remap.end()), - _flavor_factors(flavor_factors) { + _flavor_factors(flavor_factors), + _flavor_diff_xs_indices( + flavor_diff_xs_indices.begin(), flavor_diff_xs_indices.end() + ), + _flavor_subproc_indices( + flavor_subproc_indices.begin(), flavor_subproc_indices.end() + ), + _flavor_per_subproc_remap( + flavor_per_subproc_remap.begin(), flavor_per_subproc_remap.end() + ) { if (pdf_grid) { for (std::size_t i = 0; i < 2; ++i) { std::set pids; - for (auto& option : diff_xs.pid_options()) { + for (auto& option : pid_options) { pids.insert(option.at(i)); } - for (auto& option : diff_xs.pid_options()) { + for (auto& option : pid_options) { _pdf_indices.at(i).push_back( std::distance(pids.begin(), pids.find(option.at(i))) ); @@ -154,9 +202,9 @@ Integrand::Integrand( ); } _active_flavors_mask.resize(mapping.channel_count()); - std::vector mask_all(diff_xs.pid_options().size()); + std::vector mask_all(pid_options.size()); for (auto [mask, active] : zip(_active_flavors_mask, active_flavors)) { - mask.resize(diff_xs.pid_options().size()); + mask.resize(pid_options.size()); for (auto index : active) { mask.at(index) = 1.; mask_all.at(index) = 1.; @@ -181,8 +229,8 @@ std::tuple, std::vector> Integrand::latent_dims() std::vector dims{_mapping.random_dim(), 1}; std::vector is_float{true, false}; - auto flav_count = _diff_xs.pid_options().size(); - if (flav_count > 1 && !std::holds_alternative(_discrete_after)) { + auto flav_count = _pid_options.size(); + if (flav_count > 1 && !std::holds_alternative(_discrete_flavor)) { dims.push_back(1); is_float.push_back(false); if ((_pdfs.at(0) || _pdfs.at(1)) && _energy_scale) { @@ -211,7 +259,7 @@ NamedVector Integrand::compute_channel_part_ret_types() const { return Type(DataType::dt_float, acc_batch_size, {n, 4}); }; - bool has_multi_flavor = _diff_xs.pid_options().size() > 1; + bool has_multi_flavor = _pid_options.size() > 1; int particle_count = static_cast(_mapping.particle_count()); int random_dim = static_cast(_mapping.random_dim()); @@ -244,17 +292,17 @@ NamedVector Integrand::compute_channel_part_ret_types() const { ret.push_back("x2_acc", acc_float); ret.push_back("flavor_id", acc_int); ret.push_back("weight_after_cuts", acc_float); - if (_madnis_training && !std::holds_alternative(_discrete_after)) { + if (_madnis_training && !std::holds_alternative(_discrete_flavor)) { ret.push_back("extra_weight_after_cuts", acc_float); } ret.push_back("ren_scale", acc_float); if ((_pdfs.at(0) || _pdfs.at(1)) && _energy_scale) { - auto flav_count = static_cast(_diff_xs.pid_options().size()); + auto flav_count = static_cast(_pid_options.size()); if (has_multi_flavor) { ret.push_back("pdf_prior", acc_float_array(flav_count)); } for (std::size_t i = 0; i < 2; ++i) { - if (_diff_xs.has_pdf(i)) { + if (_diff_xs.at(0).has_pdf(i)) { ret.push_back(std::format("pdf{}", i + 1), acc_float); ret.push_back(std::format("fact_scale{}", i + 1), acc_float); } @@ -267,7 +315,7 @@ NamedVector Integrand::compute_channel_part_ret_types() const { NamedVector Integrand::build_channel_part( FunctionBuilder& fb, const NamedVector& args ) const { - bool has_multi_flavor = _diff_xs.pid_options().size() > 1; + bool has_multi_flavor = _pid_options.size() > 1; bool has_permutations = _mapping.channel_count() > 1; auto batch_size_val = args.at("batch_size"); @@ -293,9 +341,29 @@ NamedVector Integrand::build_channel_part( mirror_random = r_val; } + // Apply adaptive map (VEGAS or MadNIS flow) + Value latent = r; + ValueVec mapping_conditions, flow_conditions; + std::visit( + Overloaded{ + [&](std::monostate) {}, + [&](const auto& admap) { + auto admap_result = admap.build_forward(fb, {r}, {}); + latent = admap_result["data"]; + adaptive_probs.push_back(admap_result["det"]); + if (_madnis_training) { + extra_weights_before_cuts.push_back(admap_result["det"]); + } else { + weights_before_cuts.push_back(admap_result["det"]); + } + flow_conditions.push_back(latent); + } + }, + _adaptive_map + ); + // Sample channel permutation Value chan_index, chan_index_in_group; - ValueVec mapping_conditions, flow_conditions; if (has_permutations) { me_int_t opt_count = _channel_indices.size(); std::visit( @@ -305,9 +373,19 @@ NamedVector Integrand::build_channel_part( chan_index_in_group = index; weights_before_cuts.push_back(chan_det); }, - [&](const auto& discrete_before) { - auto discrete_result = - discrete_before.build_forward(fb, {chan_random}, {}); + [&](const auto& discrete_sym) { + ValueVec discrete_condition; + using TDiscrete = std::decay_t; + if constexpr (std::is_same_v) { + if (flow_conditions.size() == 1) { + discrete_condition.push_back(flow_conditions.at(0)); + } else if (flow_conditions.size() > 1) { + discrete_condition.push_back(fb.cat(flow_conditions)); + } + } + auto discrete_result = discrete_sym.build_forward( + fb, {chan_random}, discrete_condition + ); chan_index_in_group = discrete_result.at(0); if (_madnis_training) { extra_weights_before_cuts.push_back(discrete_result["det"]); @@ -320,7 +398,7 @@ NamedVector Integrand::build_channel_part( ); } }, - _discrete_before + _discrete_sym ); chan_index = fb.gather_int(chan_index_in_group, _channel_indices); mapping_conditions.push_back(chan_index_in_group); @@ -330,35 +408,6 @@ NamedVector Integrand::build_channel_part( chan_index_in_group = fb.full({static_cast(0), batch_size_val}); } - // Apply adaptive map (VEGAS or MadNIS flow) - Value latent = r; - std::visit( - Overloaded{ - [&](std::monostate) {}, - [&](const auto& admap) { - ValueVec cond; - using TAdaptive = std::decay_t; - if constexpr (std::is_same_v) { - if (flow_conditions.size() == 1) { - cond.push_back(flow_conditions.at(0)); - } else if (flow_conditions.size() > 1) { - cond.push_back(fb.cat(flow_conditions)); - } - } - auto admap_result = admap.build_forward(fb, {r}, cond); - latent = admap_result["data"]; - adaptive_probs.push_back(admap_result["det"]); - if (_madnis_training) { - extra_weights_before_cuts.push_back(admap_result["det"]); - } else { - weights_before_cuts.push_back(admap_result["det"]); - } - flow_conditions.push_back(latent); - } - }, - _adaptive_map - ); - // Apply phase space mapping auto mapping_result = _mapping.build_forward(fb, {latent}, mapping_conditions); weights_before_cuts.push_back(mapping_result["det"]); @@ -389,7 +438,7 @@ NamedVector Integrand::build_channel_part( if ((_pdfs.at(0) || _pdfs.at(1)) && _energy_scale) { ValueVec pdf_priors; for (std::size_t i = 0; i < 2; ++i) { - if (_diff_xs.has_pdf(i)) { + if (_diff_xs.at(0).has_pdf(i)) { auto pdf = _pdfs.at(i) .value() @@ -427,15 +476,15 @@ NamedVector Integrand::build_channel_part( } else { auto [index, flavor_det] = fb.sample_discrete( flavor_random_acc, - static_cast(_diff_xs.pid_options().size()) + static_cast(_pid_options.size()) ); flavor_id = index; weights_after_cuts.push_back(flavor_det); } }, - [&](const auto& discrete_after) { + [&](const auto& discrete_flavor) { ValueVec discrete_condition; - using TDiscrete = std::decay_t; + using TDiscrete = std::decay_t; if constexpr (std::is_same_v) { if (flow_conditions.size() == 1) { discrete_condition.push_back(flow_conditions.at(0)); @@ -446,7 +495,7 @@ NamedVector Integrand::build_channel_part( if (has_pdf_prior) { discrete_condition.push_back(pdf_prior); } - auto discrete_result = discrete_after.build_forward( + auto discrete_result = discrete_flavor.build_forward( fb, {flavor_random_acc}, discrete_condition ); flavor_id = discrete_result.at(0); @@ -461,7 +510,7 @@ NamedVector Integrand::build_channel_part( ); } }, - _discrete_after + _discrete_flavor ); for (auto [pdf, indices] : zip(pdf_results, _pdf_indices)) { if (pdf) { @@ -538,11 +587,11 @@ NamedVector Integrand::build_channel_part( if (has_pdf_prior) { out.push_back("pdf_prior", pdf_prior); } - if (_diff_xs.has_pdf(0)) { + if (_diff_xs.at(0).has_pdf(0)) { out.push_back("pdf1", pdf_results.at(0)); out.push_back("fact_scale1", scales.at(1)); } - if (_diff_xs.has_pdf(1)) { + if (_diff_xs.at(0).has_pdf(1)) { out.push_back("pdf2", pdf_results.at(1)); out.push_back("fact_scale2", scales.at(2)); } @@ -553,7 +602,7 @@ NamedVector Integrand::build_channel_part( NamedVector Integrand::build_common_part( FunctionBuilder& fb, const NamedVector& args ) const { - bool has_multi_flavor = _diff_xs.pid_options().size() > 1; + bool has_multi_flavor = _pid_options.size() > 1; bool has_permutations = _mapping.channel_count() > 1; bool has_pdf_prior = (_pdfs.at(0) || _pdfs.at(1)) && _energy_scale && has_multi_flavor; @@ -579,13 +628,24 @@ NamedVector Integrand::build_common_part( }; // Channel weight computation - std::size_t channel_count = _chan_weight_remap.size() > 0 ? _remapped_chan_count - : _subchan_weights - ? _subchan_weights->channel_count() - : _diff_xs.matrix_element().diagram_count(); + std::size_t channel_count = final_channel_count( + _diff_xs, + _first_chan_weight_remap, + _first_remapped_chan_count, + _second_chan_weight_remap, + _second_remapped_chan_count, + _subchan_weights + ); Value chan_weights_acc; if (channel_count > 1 && _prop_chan_weights) { chan_weights_acc = _prop_chan_weights->build_function(fb, {momenta_acc}).at(0); + if (_first_chan_weight_remap.size() > 0) { + chan_weights_acc = fb.collect_channel_weights( + chan_weights_acc, + _first_chan_weight_remap.at(0), + _first_remapped_chan_count + ); + } } // Compute running coupling @@ -600,14 +660,74 @@ NamedVector Integrand::build_common_part( xs_args.push_back(x1_acc); xs_args.push_back(x2_acc); xs_args.push_back(flavor_id); - if (_diff_xs.has_pdf(0)) { + if (_diff_xs.at(0).has_pdf(0)) { xs_args.push_back(args.at("pdf1")); } - if (_diff_xs.has_pdf(1)) { + if (_diff_xs.at(0).has_pdf(1)) { xs_args.push_back(args.at("pdf2")); } xs_args.push_back(alpha_qcd_acc); - auto dxs_vec = _diff_xs.build_function(fb, xs_args); + ValueVec dxs_vec; + Value ps_flavor_id; + Value subproc_id; + if (_diff_xs.size() == 1) { + dxs_vec = _diff_xs.at(0).build_function(fb, xs_args).values(); + ps_flavor_id = flavor_id; + if (channel_count > 1 && !_prop_chan_weights) { + chan_weights_acc = dxs_vec.at(1); + if (_first_chan_weight_remap.size() > 0) { + chan_weights_acc = fb.collect_channel_weights( + chan_weights_acc, + _first_chan_weight_remap.at(0), + _first_remapped_chan_count + ); + } + } + } else { + Value dxs_index = fb.gather_int(flavor_id, _flavor_diff_xs_indices); + ps_flavor_id = fb.gather_int(flavor_id, _flavor_per_subproc_remap); + subproc_id = fb.gather_int(flavor_id, _flavor_subproc_indices); + ValueVec split_indices = + fb.batch_split_by_index(dxs_index, static_cast(_diff_xs.size())); + std::vector split_outputs(_diff_xs.at(0).return_types().size()); + ValueVec split_channel_weights; + for (std::size_t i = 0; + auto [diff_xs, indices] : zip(_diff_xs, split_indices)) { + ValueVec split_args; + for (Value& arg : xs_args) { + split_args.push_back(fb.batch_gather(indices, arg)); + } + ValueVec outputs = diff_xs.build_function(fb, split_args).values(); + for (auto [out, split_out] : zip(outputs, split_outputs)) { + split_out.push_back(out); + split_out.push_back(indices); + } + if (channel_count > 1 && !_prop_chan_weights) { + Value split_cw = outputs.at(1); + if (_first_chan_weight_remap.size() > 0) { + split_cw = fb.collect_channel_weights( + split_cw, + _first_chan_weight_remap.at(i), + _first_remapped_chan_count + ); + } + split_channel_weights.push_back(split_cw); + split_channel_weights.push_back(indices); + } + ++i; + } + for (std::size_t i = 0; auto& split_out : split_outputs) { + if (i == 1) { + dxs_vec.push_back(Value()); + } else { + dxs_vec.push_back(fb.batch_merge_by_index(split_out)); + } + ++i; + } + if (channel_count > 1 && !_prop_chan_weights) { + chan_weights_acc = fb.batch_merge_by_index(split_channel_weights); + } + } auto diff_xs_acc = dxs_vec.at(0); if (_flavor_factors.size() > 0) { diff_xs_acc = fb.mul(diff_xs_acc, fb.gather(flavor_id, _flavor_factors)); @@ -617,22 +737,19 @@ NamedVector Integrand::build_common_part( if (args.index_map().contains("extra_weight_after_cuts")) { extra_weights_after_cuts.push_back(args.at("extra_weight_after_cuts")); } - if (!_prop_chan_weights) { - chan_weights_acc = dxs_vec.at(1); - } if (channel_count > 1 && _subchan_weights) { chan_weights_acc = _subchan_weights->build_function(fb, {momenta_acc, chan_weights_acc}).at(0); } - if (_chan_weight_remap.size() > 0) { + if (_second_chan_weight_remap.size() > 0) { chan_weights_acc = fb.collect_channel_weights( - chan_weights_acc, _chan_weight_remap, _remapped_chan_count + chan_weights_acc, _second_chan_weight_remap, _second_remapped_chan_count ); } // Apply channel weight network auto prior_chan_weights_acc = chan_weights_acc; - if (_chan_weight_net) { + if (channel_count > 1 && _chan_weight_net) { auto& preproc = _chan_weight_net.value().preprocessing(); auto cw_preproc_acc = preproc.build_function(fb, {momenta_acc, x1_acc, x2_acc}).at(0); @@ -713,13 +830,13 @@ NamedVector Integrand::build_common_part( "channel_index_in_group", optional_cut(args.at("chan_index_in_group")) ); if (has_multi_flavor && - !std::holds_alternative(_discrete_after)) { + !std::holds_alternative(_discrete_flavor)) { auto zeros = fb.full({static_cast(0), batch_size_val}); outputs.push_back( "discrete_flavor_index", scatter_or_drop(zeros, flavor_id) ); if (has_pdf_prior) { - auto flav_count = static_cast(_diff_xs.pid_options().size()); + auto flav_count = static_cast(_pid_options.size()); outputs.push_back( "pdf_prior", scatter_or_drop( @@ -744,7 +861,12 @@ NamedVector Integrand::build_common_part( outputs.push_back("color_index", scatter_or_drop(zeros_int, dxs_vec.at(2))); outputs.push_back("helicity_index", scatter_or_drop(zeros_int, dxs_vec.at(3))); outputs.push_back("diagram_index", scatter_or_drop(zeros_int, dxs_vec.at(4))); - outputs.push_back("flavor_index", scatter_or_drop(zeros_int, flavor_id)); + outputs.push_back("flavor_index", scatter_or_drop(zeros_int, ps_flavor_id)); + if (subproc_id) { + outputs.push_back( + "subprocess_index", scatter_or_drop(zeros_int, subproc_id) + ); + } outputs.push_back( "ren_scale", scatter_or_drop(zeros_float, args.at("ren_scale")) @@ -752,7 +874,7 @@ NamedVector Integrand::build_common_part( outputs.push_back("alpha_qcd", scatter_or_drop(zeros_float, alpha_qcd_acc)); if (_partial_weights) { ValueVec pdf_vals; - if (_diff_xs.has_pdf(0)) { + if (_diff_xs.at(0).has_pdf(0)) { outputs.push_back( "x1", scatter_or_drop(zeros_float, args.at("x1_acc")) ); @@ -761,7 +883,7 @@ NamedVector Integrand::build_common_part( ); pdf_vals.push_back(args.at("pdf1")); } - if (_diff_xs.has_pdf(1)) { + if (_diff_xs.at(0).has_pdf(1)) { outputs.push_back( "x2", scatter_or_drop(zeros_float, args.at("x2_acc")) ); @@ -770,7 +892,7 @@ NamedVector Integrand::build_common_part( ); pdf_vals.push_back(args.at("pdf2")); } - if (_diff_xs.has_pdf(0) || _diff_xs.has_pdf(1)) { + if (_diff_xs.at(0).has_pdf(0) || _diff_xs.at(0).has_pdf(1)) { outputs.push_back( "partial_weight_product", scatter_or_drop(zeros_float, fb.product(pdf_vals)) @@ -779,13 +901,13 @@ NamedVector Integrand::build_common_part( } outputs.push_back("random", optional_cut(args.at("r"))); if (has_permutations && - !std::holds_alternative(_discrete_before)) { + !std::holds_alternative(_discrete_sym)) { outputs.push_back( "channel_index_in_group", optional_cut(args.at("chan_index_in_group")) ); } if (has_multi_flavor && - !std::holds_alternative(_discrete_after)) { + !std::holds_alternative(_discrete_flavor)) { outputs.push_back( "discrete_flavor_index", scatter_or_drop(zeros_int, flavor_id) ); @@ -961,9 +1083,9 @@ IntegrandProbability::IntegrandProbability(const Integrand& integrand) : {"latent", batch_float_array(integrand._mapping.random_dim())}, {"channel_index_in_group", batch_int} }; - auto flavor_count = integrand._diff_xs.pid_options().size(); + auto flavor_count = integrand._pid_options.size(); if (flavor_count > 1 && - !std::holds_alternative(integrand._discrete_after)) { + !std::holds_alternative(integrand._discrete_flavor)) { arg_types.push_back("discrete_flavor_index", batch_int); if ((integrand._pdfs.at(0) || integrand._pdfs.at(1)) && integrand._energy_scale) { @@ -975,10 +1097,10 @@ IntegrandProbability::IntegrandProbability(const Integrand& integrand) : {{"prob", batch_float}} ), _adaptive_map(integrand._adaptive_map), - _discrete_before(integrand._discrete_before), - _discrete_after(integrand._discrete_after), + _discrete_sym(integrand._discrete_sym), + _discrete_flavor(integrand._discrete_flavor), _permutation_count(integrand._mapping.channel_count()), - _flavor_count(integrand._diff_xs.pid_options().size()), + _flavor_count(integrand._pid_options.size()), _has_pdf_prior( (integrand._pdfs.at(0) || integrand._pdfs.at(1)) && integrand._energy_scale ) {} @@ -987,23 +1109,6 @@ NamedVector IntegrandProbability::build_function_impl( FunctionBuilder& fb, const NamedVector& args ) const { ValueVec probs, flow_conditions; - if (_permutation_count > 1) { - auto chan_index = args.at(1); - std::visit( - Overloaded{ - [](std::monostate) {}, - [&](const auto& discrete_before) { - auto discrete_result = - discrete_before.build_inverse(fb, {chan_index}, {}); - probs.push_back(discrete_result["det"]); - flow_conditions.push_back(fb.one_hot( - chan_index, static_cast(_permutation_count) - )); - } - }, - _discrete_before - ); - } auto latent = args.at(0); std::visit( @@ -1011,15 +1116,7 @@ NamedVector IntegrandProbability::build_function_impl( [&](std::monostate) {}, [&](const auto& admap) { ValueVec cond; - using TAdaptive = std::decay_t; - if constexpr (std::is_same_v) { - if (flow_conditions.size() == 1) { - cond.push_back(flow_conditions.at(0)); - } else if (flow_conditions.size() > 1) { - cond.push_back(fb.cat(flow_conditions)); - } - } - auto admap_result = admap.build_inverse(fb, {latent}, cond); + auto admap_result = admap.build_inverse(fb, {latent}, {}); probs.push_back(admap_result["det"]); flow_conditions.push_back(latent); } @@ -1027,16 +1124,44 @@ NamedVector IntegrandProbability::build_function_impl( _adaptive_map ); + if (_permutation_count > 1) { + auto chan_index = args.at(1); + std::visit( + Overloaded{ + [](std::monostate) {}, + [&](const auto& discrete_sym) { + ValueVec discrete_condition; + using TDiscrete = std::decay_t; + if constexpr (std::is_same_v) { + if (flow_conditions.size() == 1) { + discrete_condition.push_back(flow_conditions.at(0)); + } else if (flow_conditions.size() > 1) { + discrete_condition.push_back(fb.cat(flow_conditions)); + } + } + auto discrete_result = discrete_sym.build_inverse( + fb, {chan_index}, discrete_condition + ); + probs.push_back(discrete_result["det"]); + flow_conditions.push_back(fb.one_hot( + chan_index, static_cast(_permutation_count) + )); + } + }, + _discrete_sym + ); + } + std::size_t arg_index = 2; if (_flavor_count > 1) { std::visit( Overloaded{ [&](std::monostate) {}, - [&](const auto& discrete_after) { + [&](const auto& discrete_flavor) { auto flavor = args.at(arg_index); ++arg_index; ValueVec discrete_condition; - using TDiscrete = std::decay_t; + using TDiscrete = std::decay_t; if constexpr (std::is_same_v) { if (flow_conditions.size() == 1) { discrete_condition.push_back(flow_conditions.at(0)); @@ -1050,11 +1175,11 @@ NamedVector IntegrandProbability::build_function_impl( discrete_condition.push_back(pdf_prior); } auto discrete_result = - discrete_after.build_inverse(fb, {flavor}, discrete_condition); + discrete_flavor.build_inverse(fb, {flavor}, discrete_condition); probs.push_back(discrete_result["det"]); } }, - _discrete_after + _discrete_flavor ); } diff --git a/madspace/src/python/instruction_set.hpp b/madspace/src/python/instruction_set.hpp index 73fceb63d..7c4129e8f 100644 --- a/madspace/src/python/instruction_set.hpp +++ b/madspace/src/python/instruction_set.hpp @@ -27,6 +27,8 @@ void add_instructions(py::classh& fb) { fb.def("squeeze", &FunctionBuilder::squeeze, py::arg("input")); fb.def("unsqueeze", &FunctionBuilder::unsqueeze, py::arg("input")); fb.def("accept_norm", &FunctionBuilder::accept_norm, py::arg("accepted_batch"), py::arg("full_batch")); + fb.def("batch_split_by_index", &FunctionBuilder::batch_split_by_index, py::arg("indices"), py::arg("count")); + fb.def("batch_merge_by_index", &FunctionBuilder::batch_merge_by_index, py::arg("args")); 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")); diff --git a/madspace/src/python/madspace.cpp b/madspace/src/python/madspace.cpp index 5e3ccf0e2..a988b632e 100644 --- a/madspace/src/python/madspace.cpp +++ b/madspace/src/python/madspace.cpp @@ -1215,16 +1215,19 @@ PYBIND11_MODULE(_madspace_py, m) { .def( py::init< const PhaseSpaceMapping&, - const DifferentialCrossSection&, + const std::vector&, const Integrand::AdaptiveMapping&, const Integrand::AdaptiveDiscrete&, const Integrand::AdaptiveDiscrete&, + const nested_vector2&, const std::optional&, const std::optional&, const std::optional&, const std::optional&, const std::optional&, const std::optional&, + const nested_vector2&, + std::size_t, const std::vector&, std::size_t, bool, @@ -1234,20 +1237,26 @@ PYBIND11_MODULE(_madspace_py, m) { const nested_vector2&, const std::vector&, const std::vector&, - const std::vector&>(), + const std::vector&, + const std::vector&, + const std::vector&, + const std::vector&>(), py::arg("mapping"), py::arg("diff_xs"), py::arg("adaptive_map") = std::monostate{}, - py::arg("discrete_before") = std::monostate{}, - py::arg("discrete_after") = std::monostate{}, + py::arg("discrete_sym") = std::monostate{}, + py::arg("discrete_flavor") = std::monostate{}, + py::arg("pid_options") = nested_vector2{}, py::arg("pdf_grid") = std::nullopt, py::arg("running_coupling") = std::nullopt, py::arg("energy_scale") = std::nullopt, py::arg("prop_chan_weights") = std::nullopt, py::arg("subchan_weights") = std::nullopt, py::arg("chan_weight_net") = std::nullopt, - py::arg("chan_weight_remap") = std::vector{}, - py::arg("remapped_chan_count") = 0, + py::arg("first_chan_weight_remap") = nested_vector2{}, + py::arg("first_remapped_chan_count") = 0, + py::arg("second_chan_weight_remap") = std::vector{}, + py::arg("second_remapped_chan_count") = 0, py::arg("madnis_training") = false, py::arg("drop_cuts_and_rescale") = false, py::arg("partial_weights") = false, @@ -1255,7 +1264,10 @@ PYBIND11_MODULE(_madspace_py, m) { py::arg("active_flavors") = nested_vector2{}, py::arg("flavor_remap") = std::vector{}, py::arg("flavor_factors") = std::vector{}, - py::arg("flavor_mirror") = std::vector{} + py::arg("flavor_mirror") = std::vector{}, + py::arg("flavor_diff_xs_indices") = std::vector{}, + py::arg("flavor_subproc_indices") = std::vector{}, + py::arg("flavor_per_subproc_remap") = std::vector{} ) .def("particle_count", &Integrand::particle_count) .def("madnis_training", &Integrand::madnis_training) @@ -1263,8 +1275,8 @@ PYBIND11_MODULE(_madspace_py, m) { .def("mapping", &Integrand::mapping) .def("diff_xs", &Integrand::diff_xs) .def("adaptive_map", &Integrand::adaptive_map) - .def("discrete_before", &Integrand::discrete_before) - .def("discrete_after", &Integrand::discrete_after) + .def("discrete_sym", &Integrand::discrete_sym) + .def("discrete_flavor", &Integrand::discrete_flavor) .def("energy_scale", &Integrand::energy_scale) .def("prop_chan_weights", &Integrand::prop_chan_weights) .def("chan_weight_net", &Integrand::chan_weight_net)