From 4ac7c9291f8db61ed4cc943215d11df4dfbc4aac Mon Sep 17 00:00:00 2001 From: Alex Dickson Date: Mon, 3 Aug 2026 12:51:08 -0400 Subject: [PATCH] change encoding back to float in E2EDiffConf --- openmmapi/include/PyTorchForce.h | 6 +++--- openmmapi/src/PyTorchForceE2EDiffConf.cpp | 4 ++-- platforms/cuda/src/CudaPyTorchKernelsE2EDiffConf.cpp | 4 ++-- .../cuda/tests/TestCudaPytorchForceE2EDiffConf.cpp | 2 +- .../src/ReferencePyTorchKernelsE2EDiffConf.cpp | 4 ++-- .../tests/TestReferencePyTorchForceE2EDiffConf.cpp | 2 +- python/mlforce.i | 4 ++-- serialization/src/PyTorchForceE2EDiffConfProxy.cpp | 10 +++++----- 8 files changed, 18 insertions(+), 18 deletions(-) diff --git a/openmmapi/include/PyTorchForce.h b/openmmapi/include/PyTorchForce.h index dac7dfc..b356deb 100755 --- a/openmmapi/include/PyTorchForce.h +++ b/openmmapi/include/PyTorchForce.h @@ -392,7 +392,7 @@ class PyTorchForceE2EDirect::GlobalParameterInfo { std::vector> pairs, std::vector> tetras, std::vector> cistrans, - std::vector> encoding); + std::vector> encoding); /** * Get the path to the file containing the graph. */ @@ -408,7 +408,7 @@ class PyTorchForceE2EDirect::GlobalParameterInfo { const std::vector> getPairs() const; const std::vector> getTetras() const; const std::vector> getCisTrans() const; - const std::vector> getEncoding() const; + const std::vector> getEncoding() const; void setUsesPeriodicBoundaryConditions(bool periodic); @@ -469,7 +469,7 @@ class PyTorchForceE2EDirect::GlobalParameterInfo { double scale; std::vector atoms; std::vector> bonds, angles, propers, impropers, pairs, tetras, cistrans; - std::vector> encoding; + std::vector> encoding; bool usePeriodic; std::vector globalParameters; diff --git a/openmmapi/src/PyTorchForceE2EDiffConf.cpp b/openmmapi/src/PyTorchForceE2EDiffConf.cpp index d68e899..1a23b45 100644 --- a/openmmapi/src/PyTorchForceE2EDiffConf.cpp +++ b/openmmapi/src/PyTorchForceE2EDiffConf.cpp @@ -21,7 +21,7 @@ PyTorchForceE2EDiffConf::PyTorchForceE2EDiffConf(const std::string& file, const std::vector> pairs, const std::vector> tetras, const std::vector> cistrans, - const std::vector> encoding + const std::vector> encoding ): file(file), @@ -71,7 +71,7 @@ const std::vector> PyTorchForceE2EDiffConf::getTetras() const { const std::vector> PyTorchForceE2EDiffConf::getCisTrans() const { return cistrans; } -const std::vector> PyTorchForceE2EDiffConf::getEncoding() const { +const std::vector> PyTorchForceE2EDiffConf::getEncoding() const { return encoding; } diff --git a/platforms/cuda/src/CudaPyTorchKernelsE2EDiffConf.cpp b/platforms/cuda/src/CudaPyTorchKernelsE2EDiffConf.cpp index c609803..8285edd 100644 --- a/platforms/cuda/src/CudaPyTorchKernelsE2EDiffConf.cpp +++ b/platforms/cuda/src/CudaPyTorchKernelsE2EDiffConf.cpp @@ -178,7 +178,7 @@ void CudaCalcPyTorchForceE2EDiffConfKernel::initialize(const System& system, con vector> tmpPairs = force.getPairs(); vector> tmpTetras = force.getTetras(); vector> tmpCisTrans = force.getCisTrans(); - vector> tmpEncoding = force.getEncoding(); + vector> tmpEncoding = force.getEncoding(); usePeriodic = force.usesPeriodicBoundaryConditions(); @@ -282,7 +282,7 @@ void CudaCalcPyTorchForceE2EDiffConfKernel::initialize(const System& system, con // encoding for (int i = 0; i < tmpEncoding.size(); i++) { for (int j = 0; j < tmpEncoding[i].size(); j++) { - enc_acc[i][j] = float(tmpEncoding[i][j]); + enc_acc[i][j] = tmpEncoding[i][j]; } } diff --git a/platforms/cuda/tests/TestCudaPytorchForceE2EDiffConf.cpp b/platforms/cuda/tests/TestCudaPytorchForceE2EDiffConf.cpp index d46e526..bf3b530 100644 --- a/platforms/cuda/tests/TestCudaPytorchForceE2EDiffConf.cpp +++ b/platforms/cuda/tests/TestCudaPytorchForceE2EDiffConf.cpp @@ -107,7 +107,7 @@ void testForce() { vector> cistrans = {}; // ENCODING (truncated to save space – full version can be inserted similarly) - vector> encoding = { + vector> encoding = { {-0.081649, -0.0015378, 0.10050, 0.082201, -0.047480, -0.012450, -0.010027, 0.0048477, -0.016389, 0.020203, -0.032699, 0.022310, 0.049317, -0.11148, -0.029759, diff --git a/platforms/reference/src/ReferencePyTorchKernelsE2EDiffConf.cpp b/platforms/reference/src/ReferencePyTorchKernelsE2EDiffConf.cpp index 35c95e9..4df8ba7 100644 --- a/platforms/reference/src/ReferencePyTorchKernelsE2EDiffConf.cpp +++ b/platforms/reference/src/ReferencePyTorchKernelsE2EDiffConf.cpp @@ -182,7 +182,7 @@ void ReferenceCalcPyTorchForceE2EDiffConfKernel::initialize(const System& system vector> tmpPairs = force.getPairs(); vector> tmpTetras = force.getTetras(); vector> tmpCisTrans = force.getCisTrans(); - vector> tmpEncoding = force.getEncoding(); + vector> tmpEncoding = force.getEncoding(); usePeriodic = force.usesPeriodicBoundaryConditions(); @@ -286,7 +286,7 @@ void ReferenceCalcPyTorchForceE2EDiffConfKernel::initialize(const System& system // encoding for (int i = 0; i < tmpEncoding.size(); i++) { for (int j = 0; j < tmpEncoding[i].size(); j++) { - enc_acc[i][j] = float(tmpEncoding[i][j]); + enc_acc[i][j] = tmpEncoding[i][j]; } } diff --git a/platforms/reference/tests/TestReferencePyTorchForceE2EDiffConf.cpp b/platforms/reference/tests/TestReferencePyTorchForceE2EDiffConf.cpp index 1e3a0bc..e0aeb54 100644 --- a/platforms/reference/tests/TestReferencePyTorchForceE2EDiffConf.cpp +++ b/platforms/reference/tests/TestReferencePyTorchForceE2EDiffConf.cpp @@ -107,7 +107,7 @@ void testForce() { vector> cistrans = {}; // ENCODING (truncated to save space – full version can be inserted similarly) - vector> encoding = { + vector> encoding = { {-0.081649, -0.0015378, 0.10050, 0.082201, -0.047480, -0.012450, -0.010027, 0.0048477, -0.016389, 0.020203, -0.032699, 0.022310, 0.049317, -0.11148, -0.029759, diff --git a/python/mlforce.i b/python/mlforce.i index f0ac2eb..08d9490 100755 --- a/python/mlforce.i +++ b/python/mlforce.i @@ -132,7 +132,7 @@ class PyTorchForceE2EDiffConf : public OpenMM::Force { const std::vector> pairs, const std::vector> tetras, const std::vector> cistrans, - const std::vector> encoding + const std::vector> encoding ); const std::string& getFile() const; @@ -145,7 +145,7 @@ class PyTorchForceE2EDiffConf : public OpenMM::Force { const std::vector> getPairs() const; const std::vector> getTetras() const; const std::vector> getCisTrans() const; - const std::vector> getEncoding() const; + const std::vector> getEncoding() const; const std::vector getParticleIndices() const; const std::vector getSignalForceWeights() const; diff --git a/serialization/src/PyTorchForceE2EDiffConfProxy.cpp b/serialization/src/PyTorchForceE2EDiffConfProxy.cpp index 3860ca5..145140e 100644 --- a/serialization/src/PyTorchForceE2EDiffConfProxy.cpp +++ b/serialization/src/PyTorchForceE2EDiffConfProxy.cpp @@ -41,7 +41,7 @@ void PyTorchForceE2EDiffConfProxy::serialize(const void* object, SerializationNo std::vector> pairs = force.getPairs(); std::vector> tetras = force.getTetras(); std::vector> cistrans = force.getCisTrans(); - std::vector> encoding = force.getEncoding(); + std::vector> encoding = force.getEncoding(); SerializationNode& atomTypeNode = node.createChildNode("AtomType"); for (int i = 0; i < atoms.size(); i++) { @@ -115,7 +115,7 @@ void PyTorchForceE2EDiffConfProxy::serialize(const void* object, SerializationNo for (int i = 0; i < encoding.size(); i++) { SerializationNode& indexNode = encodingNode.createChildNode("Indexes"); for (int j = 0; j < encoding[0].size(); j++) { - indexNode.createChildNode("Value").setDoubleProperty("value",encoding[i][j]); + indexNode.createChildNode("Value").setDoubleProperty("value",double(encoding[i][j])); } } } @@ -144,7 +144,7 @@ void* PyTorchForceE2EDiffConfProxy::deserialize(const SerializationNode& node) c std::vector> pairs; std::vector> tetras; std::vector> cistrans; - std::vector> encoding; + std::vector> encoding; const SerializationNode& atomTypeNode = node.getChildNode("AtomType"); for (auto &type: atomTypeNode.getChildren()) { @@ -223,9 +223,9 @@ void* PyTorchForceE2EDiffConfProxy::deserialize(const SerializationNode& node) c const SerializationNode& encodingNode = node.getChildNode("Encoding"); for (auto &indexNode: encodingNode.getChildren()) { - std::vector tmp; + std::vector tmp; for (auto &valueNode: indexNode.getChildren()) { - tmp.push_back(valueNode.getDoubleProperty("value")); + tmp.push_back(float(valueNode.getDoubleProperty("value"))); } encoding.push_back(tmp); }