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
2 changes: 2 additions & 0 deletions dwave/plugins/torch/models/boltzmann_machine.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,8 @@ def set_quadratic(self, quadratic: dict[tuple[Hashable, Hashable], float]) -> No
"""
for (u, v), bias in quadratic.items():
idx = self._edge_to_idx.get((u, v), self._edge_to_idx.get((v, u)))
if idx is None:
raise ValueError(f"Edge {(u, v)!r} is not in the model.")
self._quadratic.data[idx] = bias

def _setup_hidden(self):
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
fixes:
- |
Raise a ``ValueError`` when ``GraphRestrictedBoltzmannMachine.set_quadratic`` receives an edge
that is not in the model. Previously, an unknown edge overwrote every quadratic bias.
10 changes: 9 additions & 1 deletion tests/test_boltzmann_machine.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,10 +82,18 @@ def test_selfloop(self):
GRBM(self.nodes, self.edges, None, {"a": 0}, {("b", "c"): 0})

def test_quadratic(self):
self.bm.set_quadratic({("d", "b"): 999})
self.bm.set_quadratic({("b", "a"): 999})
self.assertEqual(999, self.bm.quadratic[0])
self.bm.set_quadratic({})

def test_set_quadratic_unknown_edge(self):
quadratic = self.bm.quadratic.detach().clone()

with self.assertRaisesRegex(ValueError, r"Edge \('d', 'b'\) is not in the model"):
self.bm.set_quadratic({("d", "b"): 999})

torch.testing.assert_close(self.bm.quadratic, quadratic)

def test_set_linear(self):
self.bm.set_linear({"d": 999})
self.assertEqual(999, self.bm.linear[0])
Expand Down