Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions openmmapi/include/PyTorchForce.h
Original file line number Diff line number Diff line change
Expand Up @@ -392,7 +392,7 @@ class PyTorchForceE2EDirect::GlobalParameterInfo {
std::vector<std::vector<int>> pairs,
std::vector<std::vector<int>> tetras,
std::vector<std::vector<int>> cistrans,
std::vector<std::vector<double>> encoding);
std::vector<std::vector<float>> encoding);
/**
* Get the path to the file containing the graph.
*/
Expand All @@ -408,7 +408,7 @@ class PyTorchForceE2EDirect::GlobalParameterInfo {
const std::vector<std::vector<int>> getPairs() const;
const std::vector<std::vector<int>> getTetras() const;
const std::vector<std::vector<int>> getCisTrans() const;
const std::vector<std::vector<double>> getEncoding() const;
const std::vector<std::vector<float>> getEncoding() const;


void setUsesPeriodicBoundaryConditions(bool periodic);
Expand Down Expand Up @@ -469,7 +469,7 @@ class PyTorchForceE2EDirect::GlobalParameterInfo {
double scale;
std::vector<int> atoms;
std::vector<std::vector<int>> bonds, angles, propers, impropers, pairs, tetras, cistrans;
std::vector<std::vector<double>> encoding;
std::vector<std::vector<float>> encoding;

bool usePeriodic;
std::vector<GlobalParameterInfo> globalParameters;
Expand Down
4 changes: 2 additions & 2 deletions openmmapi/src/PyTorchForceE2EDiffConf.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ PyTorchForceE2EDiffConf::PyTorchForceE2EDiffConf(const std::string& file,
const std::vector<std::vector<int>> pairs,
const std::vector<std::vector<int>> tetras,
const std::vector<std::vector<int>> cistrans,
const std::vector<std::vector<double>> encoding
const std::vector<std::vector<float>> encoding
):

file(file),
Expand Down Expand Up @@ -71,7 +71,7 @@ const std::vector<std::vector<int>> PyTorchForceE2EDiffConf::getTetras() const {
const std::vector<std::vector<int>> PyTorchForceE2EDiffConf::getCisTrans() const {
return cistrans;
}
const std::vector<std::vector<double>> PyTorchForceE2EDiffConf::getEncoding() const {
const std::vector<std::vector<float>> PyTorchForceE2EDiffConf::getEncoding() const {
return encoding;
}

Expand Down
4 changes: 2 additions & 2 deletions platforms/cuda/src/CudaPyTorchKernelsE2EDiffConf.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -178,7 +178,7 @@ void CudaCalcPyTorchForceE2EDiffConfKernel::initialize(const System& system, con
vector<vector<int>> tmpPairs = force.getPairs();
vector<vector<int>> tmpTetras = force.getTetras();
vector<vector<int>> tmpCisTrans = force.getCisTrans();
vector<vector<double>> tmpEncoding = force.getEncoding();
vector<vector<float>> tmpEncoding = force.getEncoding();
usePeriodic = force.usesPeriodicBoundaryConditions();


Expand Down Expand Up @@ -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];
}
}

Expand Down
2 changes: 1 addition & 1 deletion platforms/cuda/tests/TestCudaPytorchForceE2EDiffConf.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,7 @@ void testForce() {
vector<vector<int>> cistrans = {};

// ENCODING (truncated to save space – full version can be inserted similarly)
vector<vector<double>> encoding = {
vector<vector<float>> 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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -182,7 +182,7 @@ void ReferenceCalcPyTorchForceE2EDiffConfKernel::initialize(const System& system
vector<vector<int>> tmpPairs = force.getPairs();
vector<vector<int>> tmpTetras = force.getTetras();
vector<vector<int>> tmpCisTrans = force.getCisTrans();
vector<vector<double>> tmpEncoding = force.getEncoding();
vector<vector<float>> tmpEncoding = force.getEncoding();
usePeriodic = force.usesPeriodicBoundaryConditions();


Expand Down Expand Up @@ -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];
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,7 @@ void testForce() {
vector<vector<int>> cistrans = {};

// ENCODING (truncated to save space – full version can be inserted similarly)
vector<vector<double>> encoding = {
vector<vector<float>> 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,
Expand Down
4 changes: 2 additions & 2 deletions python/mlforce.i
Original file line number Diff line number Diff line change
Expand Up @@ -132,7 +132,7 @@ class PyTorchForceE2EDiffConf : public OpenMM::Force {
const std::vector<std::vector<int>> pairs,
const std::vector<std::vector<int>> tetras,
const std::vector<std::vector<int>> cistrans,
const std::vector<std::vector<double>> encoding
const std::vector<std::vector<float>> encoding
);

const std::string& getFile() const;
Expand All @@ -145,7 +145,7 @@ class PyTorchForceE2EDiffConf : public OpenMM::Force {
const std::vector<std::vector<int>> getPairs() const;
const std::vector<std::vector<int>> getTetras() const;
const std::vector<std::vector<int>> getCisTrans() const;
const std::vector<std::vector<double>> getEncoding() const;
const std::vector<std::vector<float>> getEncoding() const;

const std::vector<int> getParticleIndices() const;
const std::vector<double> getSignalForceWeights() const;
Expand Down
10 changes: 5 additions & 5 deletions serialization/src/PyTorchForceE2EDiffConfProxy.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,7 @@ void PyTorchForceE2EDiffConfProxy::serialize(const void* object, SerializationNo
std::vector<std::vector<int>> pairs = force.getPairs();
std::vector<std::vector<int>> tetras = force.getTetras();
std::vector<std::vector<int>> cistrans = force.getCisTrans();
std::vector<std::vector<double>> encoding = force.getEncoding();
std::vector<std::vector<float>> encoding = force.getEncoding();

SerializationNode& atomTypeNode = node.createChildNode("AtomType");
for (int i = 0; i < atoms.size(); i++) {
Expand Down Expand Up @@ -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]));
}
}
}
Expand Down Expand Up @@ -144,7 +144,7 @@ void* PyTorchForceE2EDiffConfProxy::deserialize(const SerializationNode& node) c
std::vector<std::vector<int>> pairs;
std::vector<std::vector<int>> tetras;
std::vector<std::vector<int>> cistrans;
std::vector<std::vector<double>> encoding;
std::vector<std::vector<float>> encoding;

const SerializationNode& atomTypeNode = node.getChildNode("AtomType");
for (auto &type: atomTypeNode.getChildren()) {
Expand Down Expand Up @@ -223,9 +223,9 @@ void* PyTorchForceE2EDiffConfProxy::deserialize(const SerializationNode& node) c

const SerializationNode& encodingNode = node.getChildNode("Encoding");
for (auto &indexNode: encodingNode.getChildren()) {
std::vector<double> tmp;
std::vector<float> tmp;
for (auto &valueNode: indexNode.getChildren()) {
tmp.push_back(valueNode.getDoubleProperty("value"));
tmp.push_back(float(valueNode.getDoubleProperty("value")));
}
encoding.push_back(tmp);
}
Expand Down