From 762ca23a702a02f65879378e52f86619b0b05de7 Mon Sep 17 00:00:00 2001 From: Dylan Kennedy Date: Mon, 27 Jul 2026 02:39:05 -0700 Subject: [PATCH 1/6] Fix pydantic issues for base, optimize, and solenoid_alignment. --- bax_algorithms/pathwise/base.py | 35 ++++++---------------------- bax_algorithms/pathwise/optimize.py | 3 +++ bax_algorithms/solenoid_alignment.py | 29 +++++++++++------------ 3 files changed, 23 insertions(+), 44 deletions(-) diff --git a/bax_algorithms/pathwise/base.py b/bax_algorithms/pathwise/base.py index 8020e3c..4893a77 100644 --- a/bax_algorithms/pathwise/base.py +++ b/bax_algorithms/pathwise/base.py @@ -1,11 +1,12 @@ # to be added to basic algorithms in Xopt from abc import abstractmethod +from difflib import Differ from bax_algorithms.pathwise.optimize import VirtualOptimizer, DifferentialEvolution from botorch.models.model import Model, ModelList from botorch.sampling.pathwise.posterior_samplers import draw_matheron_paths -from pydantic import Field -from xopt.generators.bayesian.bax.algorithms import Algorithm +from pydantic import Field, field_validator +from xopt.generators.bayesian.bax.algorithms import Algorithm, VirtualMeasurementResult, OptimizationAlgorithmResult from torch import Tensor import torch from typing import List @@ -39,49 +40,27 @@ class PathwiseOptimization(Algorithm): Get the bounds for virtual optimization. """ - name = "pathwise_optimization" - optimizer: VirtualOptimizer = Field( + name: str = Field("pathwise_optimization", frozen=True) + optimizer: DifferentialEvolution = Field( DifferentialEvolution(), description="Optimizer for virtual objective." ) - results: dict = Field( - default=None, - description="dictionary containing algorithm results", - ) observable_names_ordered: List[str] = Field( - default=None, description="names of observable models used in this algorithm", ) - @abstractmethod - def perform_virtual_measurement( - self, - model: Model, - x: Tensor, - bounds: Tensor, - n_samples: int = None, - tkwargs: dict = None, - ) -> dict: - """ - Defines how the measurement of the virtual objective should be performed. - Stores results in a dictionary. - Returned dictionary must contain key 'objective' containing the virtual objective results. - """ - return {"objective": None} - def evaluate_virtual_objective( self, model: Model, x: Tensor, bounds: Tensor, - n_samples: int = None, - tkwargs: dict = None, + n_samples: int | None = None, ) -> Tensor: """ Performs virtual measurement and extracts virtual objective value from resultant dictionary. """ measurement_result = self.perform_virtual_measurement( - model, x, bounds, n_samples, tkwargs + model, x, bounds, n_samples, ) return measurement_result.objective diff --git a/bax_algorithms/pathwise/optimize.py b/bax_algorithms/pathwise/optimize.py index cd311c1..e7b5686 100644 --- a/bax_algorithms/pathwise/optimize.py +++ b/bax_algorithms/pathwise/optimize.py @@ -31,6 +31,7 @@ def optimize( """ Minimizes virtual objective sample functions and returns optimal inputs. """ + raise NotImplementedError("This method should be implemented in subclasses.") @abstractmethod def _wrap_virtual_objective( @@ -44,6 +45,7 @@ def _wrap_virtual_objective( """ Wraps virtual objective function so inputs/outputs are suitable for optimization method. """ + raise NotImplementedError("This method should be implemented in subclasses.") @abstractmethod def _get_virtual_optimization_bounds( @@ -52,6 +54,7 @@ def _get_virtual_optimization_bounds( """ Get bounds for virtual optimization (may not be the same as bounds passed to optimizer). """ + raise NotImplementedError("This method should be implemented in subclasses.") def _get_target_function( self, diff --git a/bax_algorithms/solenoid_alignment.py b/bax_algorithms/solenoid_alignment.py index 347392c..288180e 100644 --- a/bax_algorithms/solenoid_alignment.py +++ b/bax_algorithms/solenoid_alignment.py @@ -20,20 +20,20 @@ class VirtualAlignmentMeasurementResult(VirtualMeasurementResult): class PathwiseSolenoidAlignment(PathwiseOptimization): name: str = Field("PathwiseSolenoidAlignment", frozen=True) - x_key: str = Field( + x_key: str | None= Field( None, description="key designating the centroid position in x from evaluate function", ) - y_key: str = Field( + y_key: str | None = Field( None, description="key designating the centroid poisition in y from evaluate function", ) - meas_dim: int = Field( + meas_dim: int | None= Field( None, description="index identifying the measurement quad dimension in the model", ) n_steps_measurement_param: int = Field( - 3, description="number of steps to use in the virtual measurement scans" + 5, description="number of steps to use in the virtual measurement scans" ) n_batch: PositiveInt = Field( 1, @@ -57,8 +57,8 @@ def y_idx(self) -> int: return self.observable_names_ordered.index(self.y_key) def perform_virtual_measurement( - self, model, x, bounds, tkwargs: dict = None, n_samples: int = None - ): + self, model: Model, x: Tensor, bounds: Tensor, n_samples: int | None = None + ) -> VirtualAlignmentMeasurementResult: """ inputs: model: a botorch ModelListGP @@ -80,7 +80,6 @@ def perform_virtual_measurement( model, x_tuning, bounds, - tkwargs, n_samples, ) @@ -97,8 +96,8 @@ def perform_virtual_measurement( return virtual_alignment_result def get_meas_scan_inputs( - self, x_tuning: Tensor, bounds: Tensor, tkwargs: dict = None - ): + self, x_tuning: Tensor, bounds: Tensor + ) -> Tensor: """ A function that generates the inputs for virtual emittance measurement scans at the tuning configurations specified by x_tuning. @@ -119,10 +118,9 @@ def get_meas_scan_inputs( # expand the x tensor to represent quad measurement scans # at the locations in tuning parameter space specified by X - tkwargs = tkwargs if tkwargs else {"dtype": torch.double, "device": "cpu"} x_meas = torch.linspace( - *bounds.T[self.meas_dim], self.n_steps_measurement_param, **tkwargs + *bounds.T[self.meas_dim], self.n_steps_measurement_param ) # prepare column of measurement scans coordinates @@ -147,7 +145,7 @@ def get_meas_scan_inputs( return x def evaluate_posterior_misalignment( - self, model, x_tuning, bounds, tkwargs: dict = None, n_samples: int = None + self, model: Model, x_tuning: Tensor, bounds: Tensor, n_samples: int | None = None ): """ inputs: @@ -158,10 +156,9 @@ def evaluate_posterior_misalignment( """ assert len(x_tuning.shape) in [2, 3] # x_tuning must be shape (n_tuning_configs, n_tuning_dims) or (n_samples, n_tuning_configs, ndim) - tkwargs = tkwargs if tkwargs else {"dtype": torch.double, "device": "cpu"} x = self.get_meas_scan_inputs( - x_tuning, bounds, tkwargs + x_tuning, bounds ) # result shape n_tuning_configs*n_steps x ndim centroid_position = self.evaluate_virtual_observables(model, x, n_samples) @@ -188,7 +185,7 @@ def evaluate_posterior_misalignment( return misalignment - def execute(self, model: Model, bounds: Tensor) -> Tensor: + def execute(self, model: Model, bounds: Tensor) -> OptimizationAlgorithmResult: best_tuning_inputs_list = [] best_objective_list = [] best_scan_inputs_list = [] @@ -234,7 +231,7 @@ def execute(self, model: Model, bounds: Tensor) -> Tensor: return algorithm_result - def _get_optimization_indeces(self, bounds) -> Tensor: + def _get_optimization_indeces(self, bounds: Tensor) -> Tensor: """ Get indeces specifying parameters for virtual objective optimization. """ From 7fa09042a9c8114c04feb571b9ad9be72e0c0924 Mon Sep 17 00:00:00 2001 From: Dylan Kennedy Date: Mon, 27 Jul 2026 02:43:10 -0700 Subject: [PATCH 2/6] Change algorithm name to match badger validation check. --- bax_algorithms/solenoid_alignment.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/bax_algorithms/solenoid_alignment.py b/bax_algorithms/solenoid_alignment.py index 288180e..a39d000 100644 --- a/bax_algorithms/solenoid_alignment.py +++ b/bax_algorithms/solenoid_alignment.py @@ -19,7 +19,7 @@ class VirtualAlignmentMeasurementResult(VirtualMeasurementResult): class PathwiseSolenoidAlignment(PathwiseOptimization): - name: str = Field("PathwiseSolenoidAlignment", frozen=True) + name: str = Field("pathwise_solenoid_alignment", frozen=True) x_key: str | None= Field( None, description="key designating the centroid position in x from evaluate function", @@ -28,7 +28,7 @@ class PathwiseSolenoidAlignment(PathwiseOptimization): None, description="key designating the centroid poisition in y from evaluate function", ) - meas_dim: int | None= Field( + meas_dim: int | None = Field( None, description="index identifying the measurement quad dimension in the model", ) From 0b29a3bfba437dc91231f414cc12268e0f64e749 Mon Sep 17 00:00:00 2001 From: Dylan Kennedy Date: Mon, 27 Jul 2026 12:10:28 -0700 Subject: [PATCH 3/6] Add changes to emittance algorithm for compatibility with badger pydantic_editor.py --- bax_algorithms/emittance.py | 43 +++++++++++++++++++++---------------- bax_algorithms/visualize.py | 2 +- 2 files changed, 26 insertions(+), 19 deletions(-) diff --git a/bax_algorithms/emittance.py b/bax_algorithms/emittance.py index ef241c3..75db893 100644 --- a/bax_algorithms/emittance.py +++ b/bax_algorithms/emittance.py @@ -1,5 +1,5 @@ from pydantic import Field, PositiveInt -from typing import Optional +from typing import List, Optional import torch from torch import Tensor @@ -461,38 +461,40 @@ class VirtualEmittanceMeasurementResult(VirtualMeasurementResult): class EmittanceAlgorithm(Algorithm): - x_key: str = Field( + name: str = Field("minimize_emittance", frozen=True) + x_key: str | None = Field( None, description="key designating the beamsize squared output in x from evaluate function", ) - y_key: str = Field( + y_key: str | None= Field( None, description="key designating the beamsize squared output in y from evaluate function", ) energy: float = Field(1.0, description="Beam energy in [eV]") q_len: float = Field( + 0.08, description="the longitudinal thickness of the measurement quadrupole" ) - rmat_x: Tensor = Field( - None, description="tensor shape 2x2 containing downstream rmat for x dimension" + rmat_x: List[float] | None = Field( + [1.0, 1.0, 0.0, 1.0], description="List length 4 containing downstream rmat for x dimension" ) - rmat_y: Tensor = Field( - None, description="tensor shape 2x2 containing downstream rmat for y dimension" + rmat_y: List[float] | None = Field( + [1.0, 1.0, 0.0, 1.0], description="List length 4 containing downstream rmat for y dimension" ) - twiss0_x: Tensor = Field( - None, - description="1d tensor length 2 containing design x-twiss: [beta0_x, alpha0_x] (for bmag)", + twiss0_x: List[float] | None = Field( + [1.0, 0.0], + description="List length 2 containing design x-twiss: [beta0_x, alpha0_x] (for bmag)", ) - twiss0_y: Tensor = Field( - None, - description="1d tensor length 2 containing design y-twiss: [beta0_y, alpha0_y] (for bmag)", + twiss0_y: List[float] | None = Field( + [1.0, 0.0], + description="List length 2 containing design y-twiss: [beta0_y, alpha0_y] (for bmag)", ) meas_dim: int = Field( - None, + 0, description="index identifying the measurement quad dimension in the model", ) n_steps_measurement_param: int = Field( - 3, description="number of steps to use in the virtual measurement scans" + 5, description="number of steps to use in the virtual measurement scans" ) thin_lens: bool = Field( False, @@ -502,9 +504,6 @@ class EmittanceAlgorithm(Algorithm): True, description="Whether to multiply the emit by the bmag to get virtual objective.", ) - results: dict = Field( - {}, description="Dictionary to store results from emittance calculcation" - ) maxiter_fit: int = Field( 20, description="Maximum number of iterations in nonlinear emittance fitting." ) @@ -513,6 +512,13 @@ class EmittanceAlgorithm(Algorithm): description="Whether to retain beamsize values only around the minimum from each scan.", ) + + def model_post_init(self, __context): + self.rmat_x = torch.tensor(self.rmat_x, dtype=torch.double) + self.rmat_y = torch.tensor(self.rmat_y, dtype=torch.double) + self.twiss0_x = torch.tensor(self.twiss0_x, dtype=torch.double) + self.twiss0_y = torch.tensor(self.twiss0_y, dtype=torch.double) + @property def x_idx(self) -> int: """ @@ -790,6 +796,7 @@ def _crop_quad_scans( class PathwiseMinimizeEmittance(EmittanceAlgorithm, PathwiseOptimization): + name: str = Field("pathwise_minimize_emittance", frozen=True) n_batch: PositiveInt = Field( 1, description="Number of sample batches to optimize, with each batch containing self.n_samples", diff --git a/bax_algorithms/visualize.py b/bax_algorithms/visualize.py index a1d9de1..4fafdd7 100644 --- a/bax_algorithms/visualize.py +++ b/bax_algorithms/visualize.py @@ -99,7 +99,7 @@ def visualize_virtual_measurement_result( # get virtual measurement (sample) results kwargs = kwargs if kwargs else {} measurement_result = generator.algorithm.perform_virtual_measurement( - bax_model, x, bounds, tkwargs=tkwargs, n_samples=n_samples, **kwargs + bax_model, x, bounds, n_samples=n_samples, **kwargs ).model_dump() # create figure and subplots From 29b48c0c3f4dc6fa01305491c8f493a30ee4f0ce Mon Sep 17 00:00:00 2001 From: Dylan Kennedy Date: Wed, 29 Jul 2026 12:26:26 -0700 Subject: [PATCH 4/6] Change input formatting and add field validation for emittance algorithm arguments. --- bax_algorithms/emittance.py | 69 +++++++++++++++++++++++++++---------- bax_algorithms/visualize.py | 6 ++-- 2 files changed, 53 insertions(+), 22 deletions(-) diff --git a/bax_algorithms/emittance.py b/bax_algorithms/emittance.py index 75db893..367c896 100644 --- a/bax_algorithms/emittance.py +++ b/bax_algorithms/emittance.py @@ -1,5 +1,6 @@ -from pydantic import Field, PositiveInt -from typing import List, Optional +from pydantic import Field, PositiveInt, field_validator, field_serializer +from typing import List, Optional, Any +import ast import torch from torch import Tensor @@ -466,27 +467,29 @@ class EmittanceAlgorithm(Algorithm): None, description="key designating the beamsize squared output in x from evaluate function", ) - y_key: str | None= Field( + + y_key: str | None = Field( None, description="key designating the beamsize squared output in y from evaluate function", ) energy: float = Field(1.0, description="Beam energy in [eV]") q_len: float = Field( - 0.08, - description="the longitudinal thickness of the measurement quadrupole" + 0.08, description="the longitudinal thickness of the measurement quadrupole" ) - rmat_x: List[float] | None = Field( - [1.0, 1.0, 0.0, 1.0], description="List length 4 containing downstream rmat for x dimension" + rmat_x: Tensor | None = Field( + Tensor([[1.0, 1.0], [0.0, 1.0]]), + description="2x2 Tensor containing downstream rmat for x dimension", ) - rmat_y: List[float] | None = Field( - [1.0, 1.0, 0.0, 1.0], description="List length 4 containing downstream rmat for y dimension" + rmat_y: Tensor | None = Field( + Tensor([[1.0, 1.0], [0.0, 1.0]]), + description="2x2 Tensor containing downstream rmat for y dimension", ) - twiss0_x: List[float] | None = Field( - [1.0, 0.0], + twiss0_x: Tensor | None = Field( + Tensor([1.0, 0.0]), description="List length 2 containing design x-twiss: [beta0_x, alpha0_x] (for bmag)", ) - twiss0_y: List[float] | None = Field( - [1.0, 0.0], + twiss0_y: Tensor | None = Field( + Tensor([1.0, 0.0]), description="List length 2 containing design y-twiss: [beta0_y, alpha0_y] (for bmag)", ) meas_dim: int = Field( @@ -512,12 +515,40 @@ class EmittanceAlgorithm(Algorithm): description="Whether to retain beamsize values only around the minimum from each scan.", ) - - def model_post_init(self, __context): - self.rmat_x = torch.tensor(self.rmat_x, dtype=torch.double) - self.rmat_y = torch.tensor(self.rmat_y, dtype=torch.double) - self.twiss0_x = torch.tensor(self.twiss0_x, dtype=torch.double) - self.twiss0_y = torch.tensor(self.twiss0_y, dtype=torch.double) + @field_validator("rmat_x", "rmat_y", "twiss0_x", "twiss0_y", mode="before") + @classmethod + def validate_tensors(cls, v: Any) -> Tensor: + """Accept tensors, (possibly nested) lists/tuples, or their string + representations (e.g. "[[1.0, 1.0], [0.0, 1.0]]" or "1.0, 0.0") and + convert them into a double-precision tensor.""" + if isinstance(v, Tensor): + return v + if isinstance(v, str): + stripped = v.strip() + if stripped.startswith("[") or stripped.startswith("("): + v = ast.literal_eval(stripped) + else: + v = [item for item in stripped.split(",")] + if isinstance(v, (list, tuple)): + float_list = cls._to_nested_floats(v) + return torch.tensor(float_list, dtype=torch.double) + raise ValueError(f"Cannot convert {v} to a Tensor.") + + @classmethod + def _to_nested_floats(cls, v: Any) -> Any: + """Recursively convert a (possibly nested) list/tuple to floats, + preserving the nesting structure.""" + if isinstance(v, (list, tuple)): + return [cls._to_nested_floats(item) for item in v] + return float(v) + + @field_serializer("rmat_x", "rmat_y", "twiss0_x", "twiss0_y") + def serialize_tensor(self, v: Any) -> Any: + """Serialize tensor fields to nested lists so they round-trip through + JSON/YAML (pydantic otherwise dumps unknown types as ``'torch.Tensor'``).""" + if isinstance(v, Tensor): + return v.tolist() + return v @property def x_idx(self) -> int: diff --git a/bax_algorithms/visualize.py b/bax_algorithms/visualize.py index 4fafdd7..26dea64 100644 --- a/bax_algorithms/visualize.py +++ b/bax_algorithms/visualize.py @@ -91,7 +91,7 @@ def visualize_virtual_measurement_result( f"which are not in generator.vocs.variable_names." ) tkwargs = generator.tkwargs - x = _generate_input_mesh(vocs, variable_names, reference_point, n_grid, tkwargs) + x = _generate_input_mesh(vocs, variable_names, data, reference_point, n_grid, tkwargs) # get bax observable models and bounds bax_model, bounds = get_bax_model_and_bounds(generator) @@ -219,7 +219,7 @@ def plot_bax_objective_convergence( if file_name.startswith(file_prefix) ] file_names = sorted(file_names, key=lambda x: int(x[prefix_len:-ext_len])) - file_paths = [os.path.abspath(file_name) for file_name in file_names] + file_paths = [os.path.join(directory, file_name) for file_name in file_names] results_dicts = [] aggregated_results_dict = {} @@ -278,7 +278,7 @@ def plot_bax_input_convergence( if file_name.startswith(file_prefix) ] file_names = sorted(file_names, key=lambda x: int(x[prefix_len:-ext_len])) - file_paths = [os.path.abspath(file_name) for file_name in file_names] + file_paths = [os.path.join(directory, file_name) for file_name in file_names] results_dicts = [] aggregated_results_dict = {} From 3507812fe074f26a0c7157ef8c27d8755bba7ced Mon Sep 17 00:00:00 2001 From: Dylan Kennedy Date: Fri, 31 Jul 2026 10:27:47 -0700 Subject: [PATCH 5/6] Ruff formatting. --- bax_algorithms/pathwise/base.py | 11 +++++++++-- bax_algorithms/solenoid_alignment.py | 12 +++++++----- bax_algorithms/visualize.py | 4 +++- 3 files changed, 19 insertions(+), 8 deletions(-) diff --git a/bax_algorithms/pathwise/base.py b/bax_algorithms/pathwise/base.py index 4893a77..212aa5d 100644 --- a/bax_algorithms/pathwise/base.py +++ b/bax_algorithms/pathwise/base.py @@ -6,7 +6,11 @@ from botorch.models.model import Model, ModelList from botorch.sampling.pathwise.posterior_samplers import draw_matheron_paths from pydantic import Field, field_validator -from xopt.generators.bayesian.bax.algorithms import Algorithm, VirtualMeasurementResult, OptimizationAlgorithmResult +from xopt.generators.bayesian.bax.algorithms import ( + Algorithm, + VirtualMeasurementResult, + OptimizationAlgorithmResult, +) from torch import Tensor import torch from typing import List @@ -60,7 +64,10 @@ def evaluate_virtual_objective( """ measurement_result = self.perform_virtual_measurement( - model, x, bounds, n_samples, + model, + x, + bounds, + n_samples, ) return measurement_result.objective diff --git a/bax_algorithms/solenoid_alignment.py b/bax_algorithms/solenoid_alignment.py index a39d000..42b88dc 100644 --- a/bax_algorithms/solenoid_alignment.py +++ b/bax_algorithms/solenoid_alignment.py @@ -20,7 +20,7 @@ class VirtualAlignmentMeasurementResult(VirtualMeasurementResult): class PathwiseSolenoidAlignment(PathwiseOptimization): name: str = Field("pathwise_solenoid_alignment", frozen=True) - x_key: str | None= Field( + x_key: str | None = Field( None, description="key designating the centroid position in x from evaluate function", ) @@ -95,9 +95,7 @@ def perform_virtual_measurement( return virtual_alignment_result - def get_meas_scan_inputs( - self, x_tuning: Tensor, bounds: Tensor - ) -> Tensor: + def get_meas_scan_inputs(self, x_tuning: Tensor, bounds: Tensor) -> Tensor: """ A function that generates the inputs for virtual emittance measurement scans at the tuning configurations specified by x_tuning. @@ -145,7 +143,11 @@ def get_meas_scan_inputs( return x def evaluate_posterior_misalignment( - self, model: Model, x_tuning: Tensor, bounds: Tensor, n_samples: int | None = None + self, + model: Model, + x_tuning: Tensor, + bounds: Tensor, + n_samples: int | None = None, ): """ inputs: diff --git a/bax_algorithms/visualize.py b/bax_algorithms/visualize.py index 26dea64..5a20a00 100644 --- a/bax_algorithms/visualize.py +++ b/bax_algorithms/visualize.py @@ -91,7 +91,9 @@ def visualize_virtual_measurement_result( f"which are not in generator.vocs.variable_names." ) tkwargs = generator.tkwargs - x = _generate_input_mesh(vocs, variable_names, data, reference_point, n_grid, tkwargs) + x = _generate_input_mesh( + vocs, variable_names, data, reference_point, n_grid, tkwargs + ) # get bax observable models and bounds bax_model, bounds = get_bax_model_and_bounds(generator) From a8fd8d0687e2c696d2124977410c14be764b951a Mon Sep 17 00:00:00 2001 From: Dylan Kennedy Date: Fri, 31 Jul 2026 10:36:19 -0700 Subject: [PATCH 6/6] Ruff check. --- bax_algorithms/emittance.py | 2 +- bax_algorithms/pathwise/base.py | 8 ++------ 2 files changed, 3 insertions(+), 7 deletions(-) diff --git a/bax_algorithms/emittance.py b/bax_algorithms/emittance.py index 367c896..c142746 100644 --- a/bax_algorithms/emittance.py +++ b/bax_algorithms/emittance.py @@ -1,5 +1,5 @@ from pydantic import Field, PositiveInt, field_validator, field_serializer -from typing import List, Optional, Any +from typing import Optional, Any import ast import torch from torch import Tensor diff --git a/bax_algorithms/pathwise/base.py b/bax_algorithms/pathwise/base.py index 212aa5d..3f693b5 100644 --- a/bax_algorithms/pathwise/base.py +++ b/bax_algorithms/pathwise/base.py @@ -1,15 +1,11 @@ # to be added to basic algorithms in Xopt -from abc import abstractmethod -from difflib import Differ -from bax_algorithms.pathwise.optimize import VirtualOptimizer, DifferentialEvolution +from bax_algorithms.pathwise.optimize import DifferentialEvolution from botorch.models.model import Model, ModelList from botorch.sampling.pathwise.posterior_samplers import draw_matheron_paths -from pydantic import Field, field_validator +from pydantic import Field from xopt.generators.bayesian.bax.algorithms import ( Algorithm, - VirtualMeasurementResult, - OptimizationAlgorithmResult, ) from torch import Tensor import torch